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 #include "mlir/Transforms/GreedyPatternRewriteDriver.h" 24 25 using namespace mlir; 26 using namespace mlir::linalg; 27 28 /// Implementation of fusion of generic ops and indexed_generic ops. 29 static bool areElementwiseOpsFusable(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 // Only allow fusing the producer of an input operand for now. 41 // TODO: allow fusing the producer of an output operand. 42 if (consumerIdx >= consumer.getNumInputs()) 43 return false; 44 45 // Get the consumer index map. The number of results of the consumer index 46 // map must match the number of loops of the producer. 47 AffineMap consumerIndexMap = consumer.getIndexingMap(consumerIdx); 48 if (consumerIndexMap.getNumResults() != producer.getNumLoops()) 49 return false; 50 51 // Currently support only operations with single result. 52 if (producer.getNumOutputs() != 1) 53 return false; 54 55 // Finally the index_map for the result must be invertible. For now just 56 // verify it is a permutation. 57 AffineMap producerResultIndexMap = producer.getOutputIndexingMap(0); 58 return producerResultIndexMap.isPermutation(); 59 } 60 61 /// Append to `fusedOpIndexingMapAttrs` the indexing maps for the operands of 62 /// the `producer` to use in the fused operation given the indexing map of the 63 /// result of the producer in the consumer. 64 static AffineMap getIndexingMapOfProducerOperandsInCoordinatesOfFusedOp( 65 OpOperand &producerOpOperand, AffineMap producerResultIndexMap, 66 AffineMap fusedConsumerArgIndexMap) { 67 // The indexing map in the consumer op (fusedConsumerArgIndexMap) is a map 68 // from consumer loop -> consumer arg tensor index/producer result tensor 69 // index. The fused loop is same as the consumer loop. For each producer arg 70 // the indexing map to be computed is a map from consumer loop -> producer 71 // arg tensor index. 72 // producerResultIndexMap is a map from producer loop -> tensor index. 73 // Compute the inverse to get map from tensor index -> producer loop. 74 // The inverse is a map from producer result tensor index -> producer loop. 75 AffineMap invProducerResultIndexMap = 76 inversePermutation(producerResultIndexMap); 77 assert(invProducerResultIndexMap && 78 "expected producer result indexig map to be invertible"); 79 80 LinalgOp producer = cast<LinalgOp>(producerOpOperand.getOwner()); 81 // argMap is a map from producer loop -> producer arg tensor index. 82 AffineMap argMap = 83 producer.getIndexingMap(producerOpOperand.getOperandNumber()); 84 85 // Compose argMap with invProducerResultIndexMap to get a map from 86 // producer result tensor index -> producer arg tensor index. 87 AffineMap t1 = argMap.compose(invProducerResultIndexMap); 88 89 // Compose t1 with fusedConsumerArgIndexMap gives an indexing map from 90 // consumer loop/ fused loop -> producer arg tensor index. 91 return t1.compose(fusedConsumerArgIndexMap); 92 } 93 94 /// Generate the region of the fused tensor operation. The region of the fused 95 /// op must be empty. 96 static void 97 generateFusedElementwiseOpRegion(PatternRewriter &rewriter, Operation *fusedOp, 98 LinalgOp producer, LinalgOp consumer, 99 AffineMap consumerToProducerLoopsMap, 100 unsigned consumerIdx, unsigned nloops) { 101 // Build the region of the fused op. 102 Block &producerBlock = producer->getRegion(0).front(); 103 Block &consumerBlock = consumer->getRegion(0).front(); 104 Block *fusedBlock = new Block(); 105 fusedOp->getRegion(0).push_back(fusedBlock); 106 BlockAndValueMapping mapper; 107 OpBuilder::InsertionGuard guard(rewriter); 108 rewriter.setInsertionPointToStart(fusedBlock); 109 110 // The block arguments are 111 // [index_0, index_1, ... , 112 // consumer_operand_0, ... , consumer_operand_(`consumerIdx`-1), 113 // producer_operand_0, ... , producer_operand_(n-1)], 114 // consumer_operand_(`consumerIdx`), .. consumer_operand_(m-1)] 115 // , where n is the number of producer's operand and m is the number 116 // consumer's operand. 117 // If both `numProducerIndices` and `numConsumerIndices` are zero, this is a 118 // generic op. In this case, there are no indices in block arguments. 119 unsigned numProducerIndices = isa<IndexedGenericOp>(producer.getOperation()) 120 ? producer.getNumLoops() 121 : 0; 122 unsigned numConsumerIndices = isa<IndexedGenericOp>(consumer.getOperation()) 123 ? consumer.getNumLoops() 124 : 0; 125 unsigned numFusedOpIndices = 126 (isa<IndexedGenericOp>(producer.getOperation()) || 127 isa<IndexedGenericOp>(consumer.getOperation())) 128 ? std::max(producer.getNumLoops(), consumer.getNumLoops()) 129 : 0; 130 131 // 0. Firstly, add all the indices to the block arguments. 132 for (unsigned i = 0, e = numFusedOpIndices; i < e; ++i) 133 fusedBlock->addArgument(rewriter.getIndexType()); 134 // 1. Map consumer indices to fusedBlock indices 1-1. 135 mapper.map(consumerBlock.getArguments().take_front(numConsumerIndices), 136 fusedBlock->getArguments().take_front(numConsumerIndices)); 137 // 2a. Embed producer indices into fusedBlock index space 1-1. 138 for (auto it : 139 llvm::zip(producerBlock.getArguments().take_front(numProducerIndices), 140 fusedBlock->getArguments().take_front(numProducerIndices))) { 141 auto newIndex = rewriter.create<mlir::AffineApplyOp>( 142 producer.getLoc(), 143 consumerToProducerLoopsMap.getSubMap(std::get<0>(it).getArgNumber()), 144 fusedBlock->getArguments().take_front(numFusedOpIndices)); 145 mapper.map(std::get<0>(it), newIndex); 146 } 147 // 2b. Replace the producer index operations by index operations placed in the 148 // fused block using the `consumerToProducerLoopsMap` to map the index spaces. 149 unsigned numFusedOpLoops = 150 std::max(producer.getNumLoops(), consumer.getNumLoops()); 151 if (producer.hasIndexSemantics()) { 152 SmallVector<Value> fusedIndices; 153 fusedIndices.reserve(numFusedOpLoops); 154 llvm::transform(llvm::seq<int64_t>(0, numFusedOpLoops), 155 std::back_inserter(fusedIndices), [&](int64_t dim) { 156 return rewriter.create<IndexOp>(producer.getLoc(), dim); 157 }); 158 for (IndexOp indexOp : 159 llvm::make_early_inc_range(producerBlock.getOps<IndexOp>())) { 160 Value newIndex = rewriter.create<mlir::AffineApplyOp>( 161 producer.getLoc(), 162 consumerToProducerLoopsMap.getSubMap(indexOp.dim()), fusedIndices); 163 // Replace the producer index operation by the index value computed in the 164 // fused block. All remaining operations in the producer block are later 165 // on cloned to the fused block. 166 rewriter.replaceOp(indexOp, newIndex); 167 } 168 } 169 // TODO: allow fusing the producer of an output operand. 170 assert(consumerIdx < consumer.getNumInputs() && 171 "expected producer of input operand"); 172 // 3. Consumer input operands up to consumerIdx (exclusive). 173 for (BlockArgument bbArg : consumerBlock.getArguments() 174 .drop_front(numConsumerIndices) 175 .take_front(consumerIdx)) // input assumption. 176 mapper.map(bbArg, fusedBlock->addArgument(bbArg.getType())); 177 178 // Replacing consumerIdx requires getting the cloned, yielded, value from 179 // the (cloned) producer block. This happens in step 9. 180 181 // 4. Splice in producer's input operands. 182 for (BlockArgument bbArg : producerBlock.getArguments() 183 .drop_front(numProducerIndices) 184 .take_front(producer.getNumInputs())) 185 mapper.map(bbArg, fusedBlock->addArgument(bbArg.getType())); 186 187 // 4.b. Producer output operand/map that is fused needs to be mapped to the 188 // producer bbArg if it is an "initTensor" (i.e. its value is actually read). 189 assert(producer->getNumResults() == 1 && "expected single result producer"); 190 if (producer.isInitTensor(&producer.getOutputOpOperands()[0])) { 191 BlockArgument bbArg = 192 producerBlock.getArguments() 193 .drop_front(numConsumerIndices + producer.getNumInputs()) 194 // TODO: bbArg index of 195 .front(); 196 mapper.map(bbArg, fusedBlock->addArgument(bbArg.getType())); 197 } 198 // 5. Remaining consumer's input operands (drop past index `consumerIdx`). 199 for (BlockArgument bbArg : consumerBlock.getArguments() 200 .drop_front(numConsumerIndices) 201 .take_front(consumer.getNumInputs()) 202 .drop_front(consumerIdx + 1)) 203 mapper.map(bbArg, fusedBlock->addArgument(bbArg.getType())); 204 // 6. All of consumer's output operands. 205 for (BlockArgument bbArg : 206 consumerBlock.getArguments().take_back(consumer.getNumOutputs())) 207 mapper.map(bbArg, fusedBlock->addArgument(bbArg.getType())); 208 // 7. All of producer's output operands except the one fused. 209 // TODO: allow fusion of multi-result producers. 210 assert(producer->getNumResults() == 1 && "expected single result producer"); 211 212 // 8. Clone operations from producer (except the yield operation) to the fused 213 // op. 214 for (auto &op : producerBlock.without_terminator()) 215 rewriter.clone(op, mapper); 216 // 9. Now we can map the consumerBlock's `consumerIdx` block argument. Just 217 // forward the yield operand. 218 auto yieldOp = cast<linalg::YieldOp>(producerBlock.getTerminator()); 219 // TODO: allow fusion of multi-result producers. 220 assert(producer->getNumResults() == 1 && "expected single result producer"); 221 unsigned producerResultNumber = 0; 222 Value replacement = 223 mapper.lookupOrDefault(yieldOp.getOperand(producerResultNumber)); 224 // Sanity checks, if replacement is not already in the mapper then it must be 225 // produced outside. 226 if (replacement == yieldOp.getOperand(producerResultNumber)) { 227 if (auto bb = replacement.dyn_cast<BlockArgument>()) 228 assert(bb.getOwner() != &producerBlock && 229 "yielded block argument must have been mapped"); 230 else 231 assert(!producer->isAncestor(replacement.getDefiningOp()) && 232 "yielded value must have been mapped"); 233 } 234 mapper.map(consumerBlock.getArgument(consumerIdx + numConsumerIndices), 235 replacement); 236 // 10. Clone operations from the consumer to the fused op. 237 for (auto &op : consumerBlock.getOperations()) 238 rewriter.clone(op, mapper); 239 240 // Sanity checks. 241 assert(fusedBlock->getNumArguments() == 242 fusedOp->getNumOperands() + numFusedOpIndices && 243 "Ill-formed LinalgOp region"); 244 } 245 246 static Optional<SmallVector<Value, 1>> 247 fuseElementwiseOpsImpl(LinalgOp producer, OpOperand &consumerOpOperand, 248 const ControlElementwiseOpsFusionFn &controlFn, 249 PatternRewriter &rewriter) { 250 LinalgOp consumer = cast<LinalgOp>(consumerOpOperand.getOwner()); 251 unsigned consumerIdx = consumerOpOperand.getOperandNumber(); 252 if (!areElementwiseOpsFusable(producer, consumer, consumerIdx) || 253 !controlFn(producer->getResult(0), consumerOpOperand)) 254 return llvm::None; 255 256 // TODO: allow fusing the producer of an output operand. 257 assert(consumerIdx < consumer.getNumInputs() && 258 "expected producer of input operand"); 259 260 // Compute the fused operands list and indexing maps. 261 SmallVector<Value> fusedOperands; 262 SmallVector<AffineMap> fusedIndexMaps; 263 fusedOperands.reserve(producer->getNumOperands() + 264 consumer->getNumOperands()); 265 fusedIndexMaps.reserve(producer->getNumOperands() + 266 consumer->getNumOperands()); 267 // In the following, numbering matches that of `generateFusedTensorOpRegion`. 268 // 3. Consumer input operands/maps up to consumerIdx (exclusive). 269 llvm::append_range(fusedOperands, 270 consumer.getInputs().take_front(consumerIdx)); 271 llvm::append_range( 272 fusedIndexMaps, 273 ArrayRef<AffineMap>{consumer.getInputIndexingMaps()}.take_front( 274 consumerIdx)); 275 // 4. Splice in producer's input operands/maps. 276 llvm::append_range(fusedOperands, producer.getInputs()); 277 assert(producer->getNumResults() == 1 && "expected single result producer"); 278 AffineMap producerResultIndexMap = producer.getOutputIndexingMap(0); 279 for (auto &inputOpOperand : producer.getInputOpOperands()) { 280 // Compute indexing maps for the producer args in the fused operation. 281 AffineMap map = getIndexingMapOfProducerOperandsInCoordinatesOfFusedOp( 282 inputOpOperand, producerResultIndexMap, 283 consumer.getInputIndexingMap(consumerIdx)); 284 fusedIndexMaps.push_back(map); 285 } 286 // 4.b. Producer output operand/map that is fused needs to be passed if it is 287 // an "initTensor" (i.e. its value is actually read). 288 assert(producer->getNumResults() == 1 && "expected single result producer"); 289 if (producer.isInitTensor(&producer.getOutputOpOperands()[0])) { 290 llvm::append_range(fusedOperands, producer.getOutputs().take_front()); 291 // Compute indexing maps for the producer args in the fused operation. 292 AffineMap map = getIndexingMapOfProducerOperandsInCoordinatesOfFusedOp( 293 producer.getOutputOpOperands().front(), producerResultIndexMap, 294 consumer.getOutputIndexingMap(0)); 295 fusedIndexMaps.push_back(map); 296 } 297 // 5. Remaining consumer's input operands/maps (drop past index 298 // `consumerIdx`). 299 llvm::append_range(fusedOperands, 300 consumer.getInputs().drop_front(consumerIdx + 1)); 301 llvm::append_range( 302 fusedIndexMaps, 303 ArrayRef<AffineMap>{consumer.getInputIndexingMaps()}.drop_front( 304 consumerIdx + 1)); 305 // 6. All of consumer's output operands (skip operands: added by the builder). 306 // llvm::append_range(fusedOperands, consumer.getOutputs()); 307 llvm::append_range(fusedIndexMaps, consumer.getOutputIndexingMaps()); 308 // 7. All of producer's output operands/maps except the one fused. 309 // TODO: allow fusion of multi-result producers. 310 assert(producer->getNumResults() == 1 && "expected single result producer"); 311 312 // Generate the fused op. 313 Operation *fusedOp; 314 if (isa<GenericOp>(producer.getOperation()) && 315 isa<GenericOp>(consumer.getOperation())) { 316 fusedOp = rewriter.create<GenericOp>( 317 consumer.getLoc(), consumer->getResultTypes(), 318 /*inputs=*/fusedOperands, 319 // TODO: handle outputs. 320 consumer.getOutputs(), rewriter.getAffineMapArrayAttr(fusedIndexMaps), 321 consumer.iterator_types(), 322 /*doc=*/nullptr, 323 /*library_call=*/nullptr, 324 /*sparse=*/nullptr); 325 } else { 326 fusedOp = rewriter.create<IndexedGenericOp>( 327 consumer.getLoc(), consumer->getResultTypes(), 328 /*inputs=*/fusedOperands, 329 // TODO: handle outputs. 330 consumer.getOutputs(), rewriter.getAffineMapArrayAttr(fusedIndexMaps), 331 consumer.iterator_types(), 332 /*doc=*/nullptr, 333 /*library_call=*/nullptr, 334 /*sparse=*/nullptr); 335 } 336 337 // Construct an AffineMap from consumer loops to producer loops. 338 // consumer loop -> tensor index 339 AffineMap consumerResultIndexMap = consumer.getInputIndexingMap(consumerIdx); 340 // tensor index -> producer loop 341 AffineMap invProducerResultIndexMap = 342 inversePermutation(producerResultIndexMap); 343 assert(invProducerResultIndexMap && 344 "expected producer result indexig map to be invertible"); 345 // consumer loop -> producer loop 346 AffineMap consumerToProducerLoopsMap = 347 invProducerResultIndexMap.compose(consumerResultIndexMap); 348 349 generateFusedElementwiseOpRegion(rewriter, fusedOp, producer, consumer, 350 consumerToProducerLoopsMap, consumerIdx, 351 consumer.getNumLoops()); 352 return SmallVector<Value, 1>(fusedOp->getResults()); 353 } 354 355 /// Linearize the expressions in `sourceMap` based on the `reassociationMaps` 356 /// provided, given the shape of the source tensor that corresponds to the 357 /// `sourceMap`. Note that this implicitly assumes that the tensors dimensions 358 /// are "row-major" ordered logically. 359 /// 360 /// For example: 361 /// 362 /// %0 = op ... : tensor<?x?x4x5xf32> 363 /// with output index_map `affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)>` 364 /// 365 /// and reshape: 366 /// %1 = linalg.tensor_reshape %0 [affine_map<(i, j, k, l) -> (i)>, 367 /// affine_map<(i, j, k, l) -> (j, k, l)>] : 368 /// tensor<?x?x4x5xf32> into tensor<?x?xf32> 369 /// 370 /// would be rewritten into: 371 /// %0 = op ... : tensor<?x?x4x5xf32> 372 /// with output index_map 373 /// `affine_map<(d0, d1, d2, d3) -> (d0, d1 * 20 + d2 * 5 + d3)>` 374 static AffineMap linearizeCollapsedDims(AffineMap sourceMap, 375 ArrayRef<int64_t> sourceShape, 376 ArrayRef<AffineMap> reassociationMaps) { 377 SmallVector<AffineExpr, 4> resultExprs; 378 resultExprs.reserve(reassociationMaps.size()); 379 ArrayRef<AffineExpr> sourceExprs = sourceMap.getResults(); 380 MLIRContext *context = sourceMap.getContext(); 381 382 // Compute the result exprs based on the reassociation maps. 383 for (AffineMap map : reassociationMaps) { 384 ArrayRef<AffineExpr> collapsedDims = map.getResults(); 385 // Assume that they are in-order and contiguous (already checked in 386 // verifier). 387 assert(!collapsedDims.empty()); 388 unsigned startDim = 389 collapsedDims.front().cast<AffineDimExpr>().getPosition(); 390 SmallVector<int64_t, 4> sizes; 391 SmallVector<AffineExpr, 4> dimExprs; 392 for (auto en : 393 llvm::zip(sourceShape.slice(startDim, collapsedDims.size()), 394 sourceExprs.slice(startDim, collapsedDims.size()))) { 395 if (std::get<0>(en) == 1) 396 continue; 397 sizes.push_back(std::get<0>(en)); 398 dimExprs.push_back(std::get<1>(en)); 399 } 400 AffineExpr linearizedExpr = 401 makeCanonicalStridedLayoutExpr(sizes, dimExprs, context); 402 resultExprs.push_back(linearizedExpr); 403 } 404 return AffineMap::get(sourceMap.getNumDims(), sourceMap.getNumSymbols(), 405 resultExprs, context); 406 } 407 408 /// Checks if the `reshapeOp` can be fused with it consumer (if `asProducer` is 409 /// true) or its producer (if `asProducer` is false) given the indexing map at 410 /// its use. 411 static bool isTensorReshapeOpFoldableByLinearization(TensorReshapeOp reshapeOp, 412 AffineMap useIndexMap, 413 bool asProducer) { 414 RankedTensorType returnType = reshapeOp.getResultType(); 415 RankedTensorType operandType = reshapeOp.getSrcType(); 416 // Reshape is fusable with its consumer (i.e. reshape as a producer) when its 417 // operand is of lesser rank than the result. Fusing when operand has higher 418 // rank will require use of mods and divs in the indexing maps of the fused op 419 // which would make it non-invertible. Similarly reshape is fused with its 420 // producer (i.e. reshape as consumer) only if the return type has lesser 421 // rank. 422 if ((asProducer && reshapeOp.getSrcType().hasStaticShape() && 423 returnType.getRank() < operandType.getRank()) || 424 (!asProducer && reshapeOp.getResultType().hasStaticShape() && 425 operandType.getRank() < returnType.getRank())) 426 return false; 427 return useIndexMap.isPermutation(); 428 } 429 430 /// Based on the type of `op` create a linalg op of the same type, i.e. if `op` 431 /// is a linalg.generic operation, the create a `linalg.generic` operation with 432 /// the given `args`. Expects `op` to be `linalg.generic` or 433 /// `linalg.indexed_generic`. 434 template <typename... Args> 435 static LinalgOp createLinalgOpOfSameType(LinalgOp op, PatternRewriter &rewriter, 436 Args... args) { 437 if (isa<GenericOp>(op.getOperation())) 438 return rewriter.create<GenericOp>(args...); 439 if (isa<IndexedGenericOp>(op.getOperation())) 440 return rewriter.create<IndexedGenericOp>(args...); 441 llvm_unreachable( 442 "expected only linalg.generic or linalg.indexed_generic ops"); 443 return nullptr; 444 } 445 446 /// Check if the reshape operation is only expansion into/collapsing of 447 /// unit-dimension. 448 static bool isUnitDimExpansionOnly(ArrayRef<int64_t> expandedShape, 449 ArrayRef<AffineMap> reassociation) { 450 for (auto &map : reassociation) { 451 unsigned numUnitDims = 0; 452 for (AffineExpr expr : map.getResults()) { 453 unsigned position = expr.cast<AffineDimExpr>().getPosition(); 454 if (expandedShape[position] == 1) 455 numUnitDims++; 456 } 457 if (numUnitDims != map.getNumResults() - 1) 458 return false; 459 } 460 return true; 461 } 462 463 /// Conditions for folding a generic/indexed-generic operation with a reshape op 464 /// by expanding the iteration space dimensionality for tensor operations. These 465 /// are preconditions assumed by `foldReshapeByDimExpansion` which implements 466 /// the following fusion pattern. 467 /// 468 /// Consider 469 /// 470 /// %c = linalg.generic ins(%a, %b : memref<?x?x?xf32>, memref<?x?xf32>) 471 /// indexing_maps = [affine_map<(d0, d1, d2) -> (d1, d0, d2)>, 472 /// affine_map<(d0, d1, d2) -> (d1, d2)>, 473 /// affine_map<(d0, d1, d2) -> (d0, d2, d1)>] 474 /// %d = linalg.tensor_reshape %c 475 /// [affine_map<(d0, d1, d2, d3, d4, d5) -> (d0, d1)>, 476 /// affine_map<(d0, d1, d2, d3, d4, d5) -> (d2)>, 477 /// affine_map<(d0, d1, d2, d3, d4, d5) -> (d3, d4, d5)>] 478 /// : tensor<?x?x?xf32> into tensor<?x?x?x?x?x?xf32> 479 /// 480 /// The reshape can be folded into the `linalgOp` if the 481 /// generic/indexed-generic op loop dimensionality is increased to match the 482 /// result (operand) of the tensor_reshape when the reshape is expanding 483 /// (folding). The indexing_map of the fused tensor in the `linalgOp` and the 484 /// reassociation map helps compute the indexing maps of the modified op. For 485 /// the above example, based on the reassociation map it can be concluded that 486 /// 487 /// - The loop used to access the first dimension of the fused tensor is split 488 /// into two. 489 /// - The loop used to access the second dimension of the fused tensor is kept 490 /// as is. 491 /// - The loop used to access the third dimension of the fused tensor is split 492 /// into three. 493 /// 494 /// i.e. (e0, e1, e2, e3, e4) is the domain of the indexing map of the modified 495 /// op, then 496 /// 497 /// d0 -> e0, e1 498 /// d1 -> e2, e3, e4 499 /// d2 -> e5 500 /// 501 /// substituting this, the generic op can be rewritten as 502 /// 503 /// %d = linalg.generic ins(%0, %1 : ) 504 /// indexing_maps = 505 /// [affine_map<(e0, e1, e2, e3, e4, e5) -> (e2, e3, e4, e0, e1, e5)>, 506 /// affine_map<(e0, e1, e2, e3, e4, e5) -> (e2, e3, e4, e5)>, 507 /// affine_map<(e0, e1, e2, e3, e4, e5) -> (e0, e1, e5, e2, e3, e4)>] 508 /// 509 /// Since operands to the linalg generic are now 5D, reshapes can be introduced 510 /// to make it consistent 511 /// 512 /// %0 = linalg.tensor_reshape %a 513 /// [affine_map<(e0, e1, e2, e3, e4, e5) -> (e0, e1, e2), 514 /// affine_map<(e0, e1, e2, e3, e4, e5) -> (e3, e4), 515 /// affine_map<(e0, e1, e2, e3, e4, e5) -> (e5)] 516 /// : tensor<?x?x?xf32> into tensor<?x?x?x?x?x?xf32> 517 /// %1 = linalg.tensor_reshape %b 518 /// [affine_map<(e0, e1, e2, e3) -> (e0, e1, e2), 519 /// affine_map<(e0, e1, e2, e3) -> (e3)] 520 /// : tensor<?x?x?xf32> into tensor<?x?x?x?xf32> 521 /// 522 /// The added reshapes are again expanding patterns, so they will get fused 523 /// with its producers if possible. 524 static bool isFusableWithReshapeByDimExpansion(LinalgOp linalgOp, 525 unsigned fusedTensorIndex) { 526 // Is fusable only if: 527 // - The linalgOp is a generic op, or an indexed_generic. 528 // - All the indexing maps for operands and results in linalgOp are projected 529 // permutations. 530 // - The fused tensor is not a scalar. 531 // - All the loops in linalgOp are parallel loops. 532 return isa<GenericOp, IndexedGenericOp>(linalgOp.getOperation()) && 533 linalgOp.hasTensorSemantics() && 534 llvm::all_of(linalgOp.indexing_maps().getValue(), 535 [](Attribute attr) { 536 return attr.cast<AffineMapAttr>() 537 .getValue() 538 .isProjectedPermutation(); 539 }) && 540 linalgOp.getIndexingMap(fusedTensorIndex).getNumResults() > 0 && 541 llvm::all_of(linalgOp.iterator_types(), [](Attribute attr) { 542 return attr.cast<StringAttr>().getValue() == 543 getParallelIteratorTypeName(); 544 }); 545 } 546 547 namespace { 548 /// Information needed to expand a generic/indexed_generic operation to fold the 549 /// reshape with it. 550 class ExpansionInfo { 551 public: 552 // Computes the mapping from original dimensions of the op to the dimensions 553 // of the expanded op given the `indexingMap` of the fused operand/result of 554 // the generic/indexed_generic op, the `reassocationMaps` of the reshape op 555 // and the shape of the expanded op. 556 LogicalResult compute(LinalgOp linalgOp, unsigned fusedTensorIndex, 557 ArrayRef<AffineMap> reassociationMaps, 558 ArrayRef<int64_t> expandedShape); 559 unsigned getOrigOpNumDims() const { return reassociation.size(); } 560 unsigned getExpandedOpNumDims() const { return expandedOpNumDims; } 561 ReassociationIndicesRef getExpandedDims(unsigned i) const { 562 return reassociation[i]; 563 } 564 ArrayRef<int64_t> getExpandedShapeOfDim(unsigned i) const { 565 return expandedShapeMap[i]; 566 } 567 568 private: 569 /// Reassociation from the dimensions in the original operation to the 570 /// dimension of the expanded operation. 571 SmallVector<ReassociationIndices, 4> reassociation; 572 /// Mapping from extent of loops in the original operation, to the extent of 573 /// loops in the expanded operation. 574 SmallVector<SmallVector<int64_t, 4>, 4> expandedShapeMap; 575 unsigned expandedOpNumDims; 576 }; 577 } // namespace 578 579 LogicalResult ExpansionInfo::compute(LinalgOp linalgOp, 580 unsigned fusedTensorIndex, 581 ArrayRef<AffineMap> reassociationMaps, 582 ArrayRef<int64_t> expandedShape) { 583 if (reassociationMaps.empty()) 584 return failure(); 585 AffineMap fusedIndexMap = linalgOp.getIndexingMap(fusedTensorIndex); 586 587 Optional<SmallVector<int64_t, 4>> originalLoopRange = 588 linalgOp.getStaticLoopRanges(); 589 if (!originalLoopRange) 590 return linalgOp.emitError("unable to find loop range for operation"); 591 592 reassociation.clear(); 593 expandedShapeMap.clear(); 594 // Compute the number of dimension in the expanded op that correspond to each 595 // dimension of the original op. 596 SmallVector<unsigned, 4> numExpandedDims(fusedIndexMap.getNumDims(), 1); 597 expandedShapeMap.resize(fusedIndexMap.getNumDims()); 598 for (auto resultExpr : llvm::enumerate(fusedIndexMap.getResults())) { 599 unsigned pos = resultExpr.value().cast<AffineDimExpr>().getPosition(); 600 AffineMap foldedDims = reassociationMaps[resultExpr.index()]; 601 numExpandedDims[pos] = foldedDims.getNumResults(); 602 ArrayRef<int64_t> shape = 603 expandedShape.slice(foldedDims.getDimPosition(0), numExpandedDims[pos]); 604 expandedShapeMap[pos].assign(shape.begin(), shape.end()); 605 } 606 // The remaining dimensions remain the same. 607 for (unsigned i : llvm::seq<unsigned>(0, fusedIndexMap.getNumDims())) 608 if (expandedShapeMap[i].empty()) 609 expandedShapeMap[i] = {(*originalLoopRange)[i]}; 610 611 // Compute reassociation map from the original op to the expanded op. 612 unsigned sum = 0; 613 reassociation.reserve(fusedIndexMap.getNumDims()); 614 for (auto numFoldedDim : llvm::enumerate(numExpandedDims)) { 615 auto seq = llvm::seq<int64_t>(sum, sum + numFoldedDim.value()); 616 reassociation.emplace_back(seq.begin(), seq.end()); 617 sum += numFoldedDim.value(); 618 } 619 expandedOpNumDims = sum; 620 return success(); 621 } 622 623 /// Epanding the body of a linalg operation requires adaptations of the accessed 624 /// loop indices. Specifically, access of indices in the original operation need 625 /// to be replaced with linearizations of indices in the expanded op. That 626 /// requires the shape of the expanded dimensions to be static (at least all but 627 /// the most significant). For now check that these are all statically sized. 628 /// Note that this could be extended to handle dynamic case, but the 629 /// implementation below uses `affine.apply` which seems to have issues when the 630 /// shapes are not static. 631 LogicalResult isIndexedOpExpandable(LinalgOp linalgOp, 632 const ExpansionInfo &expansionInfo) { 633 for (unsigned i : llvm::seq<unsigned>(0, expansionInfo.getOrigOpNumDims())) { 634 ArrayRef<int64_t> expandedShape = expansionInfo.getExpandedShapeOfDim(i); 635 if (expandedShape.size() == 1) 636 continue; 637 for (int64_t shape : expandedShape.drop_front()) { 638 if (ShapedType::isDynamic(shape)) { 639 return linalgOp.emitError( 640 "unable to fuse indexed generic op where the expanded dim is " 641 "dynamic"); 642 } 643 } 644 } 645 return success(); 646 } 647 648 /// Return the indexing map to use in the expanded op for a given the 649 /// `indexingMap` of the original operation. 650 static AffineMap 651 getIndexingMapInExpandedOp(OpBuilder &builder, AffineMap indexingMap, 652 const ExpansionInfo &expansionInfo) { 653 SmallVector<AffineExpr, 4> newExprs; 654 for (AffineExpr expr : indexingMap.getResults()) { 655 unsigned pos = expr.cast<AffineDimExpr>().getPosition(); 656 SmallVector<AffineExpr, 4> expandedExprs = llvm::to_vector<4>( 657 llvm::map_range(expansionInfo.getExpandedDims(pos), [&](int64_t v) { 658 return builder.getAffineDimExpr(static_cast<unsigned>(v)); 659 })); 660 newExprs.append(expandedExprs.begin(), expandedExprs.end()); 661 } 662 return AffineMap::get(expansionInfo.getExpandedOpNumDims(), 663 indexingMap.getNumSymbols(), newExprs, 664 builder.getContext()); 665 } 666 667 /// Return the type of the operand/result to use in the expanded op given the 668 /// type in the original op. 669 static RankedTensorType getExpandedType(RankedTensorType originalType, 670 AffineMap indexingMap, 671 const ExpansionInfo &expansionInfo) { 672 SmallVector<int64_t, 4> expandedShape; 673 for (AffineExpr expr : indexingMap.getResults()) { 674 unsigned dim = expr.cast<AffineDimExpr>().getPosition(); 675 auto dimExpansion = expansionInfo.getExpandedShapeOfDim(dim); 676 expandedShape.append(dimExpansion.begin(), dimExpansion.end()); 677 } 678 return RankedTensorType::get(expandedShape, originalType.getElementType()); 679 } 680 681 /// Returns the reassociation maps to use in the `linalg.tensor_reshape` 682 /// operation to convert the operands of the origial operation to operands of 683 /// the expanded operation. The same method is used to compute the 684 /// `linalg.tensor_reshape` used to collapse the result of the expanded op to 685 /// get the value that can replace all uses of the results of the original op. 686 static SmallVector<ReassociationIndices, 4> 687 getReassociationForExpansion(AffineMap indexingMap, 688 const ExpansionInfo &expansionInfo) { 689 SmallVector<ReassociationIndices, 4> reassociation; 690 unsigned numReshapeDims = 0; 691 for (AffineExpr expr : indexingMap.getResults()) { 692 unsigned dim = expr.cast<AffineDimExpr>().getPosition(); 693 auto numExpandedDims = expansionInfo.getExpandedDims(dim).size(); 694 auto indices = llvm::to_vector<2>( 695 llvm::seq<int64_t>(numReshapeDims, numReshapeDims + numExpandedDims)); 696 reassociation.emplace_back(std::move(indices)); 697 numReshapeDims += numExpandedDims; 698 } 699 return reassociation; 700 } 701 702 /// Build the body of the expanded IndexedGenericOp. The arguments for the 703 /// induction variables of the original operation need to be recovered by 704 /// linearizing the arguments of the corresponding dimensions of the expanded 705 /// op. For now it is assumed that the shapes of the expanded op needed for 706 /// linearization are static. 707 static void buildExpandedIndexedGenericOpRegion( 708 PatternRewriter &rewriter, Location loc, Region &originalOpRegion, 709 Region &fusedOpRegion, const ExpansionInfo &expansionInfo) { 710 assert(fusedOpRegion.empty() && "expected fused op to have empty region"); 711 // Create an entry block in the fused region with same number of arguments 712 // as the fused op 713 Block *fusedEntryBlock = new Block; 714 fusedOpRegion.push_back(fusedEntryBlock); 715 rewriter.cloneRegionBefore(originalOpRegion, fusedOpRegion, 716 fusedOpRegion.end()); 717 718 // Merge the entry block of the fused op with the cloned blocks. For this 719 // compute the value for arguments of the region in the original operation 720 // in terms of the arguments of the fused op. Since the original operation 721 // is expanded, the expanded dimensions need to be folded back to get the 722 // replacement value for the arguments corresponding to interation index. 723 // For now this expects that all the loop ranges are constants, which is 724 // true if the shapes are all static. This has already been checked in the 725 // precondition. 726 using namespace edsc::op; 727 using namespace edsc::intrinsics; 728 OpBuilder::InsertionGuard guard(rewriter); 729 SmallVector<Value, 4> argReplacements(originalOpRegion.getNumArguments()); 730 rewriter.setInsertionPointToStart(fusedEntryBlock); 731 edsc::ScopedContext scopedContext(rewriter, loc); 732 IndexType indexType = rewriter.getIndexType(); 733 for (auto i : llvm::seq<unsigned>(0, expansionInfo.getOrigOpNumDims())) { 734 Value linearizedIndex = fusedEntryBlock->addArgument(indexType); 735 ArrayRef<int64_t> expandedDimsShape = 736 expansionInfo.getExpandedShapeOfDim(i).drop_front(); 737 for (unsigned shape : expandedDimsShape) { 738 assert(!ShapedType::isDynamic(shape)); 739 linearizedIndex = linearizedIndex * std_constant_index(shape); 740 linearizedIndex = 741 linearizedIndex + fusedEntryBlock->addArgument(indexType); 742 } 743 argReplacements[i] = linearizedIndex; 744 } 745 for (auto i : llvm::seq<unsigned>(expansionInfo.getOrigOpNumDims(), 746 argReplacements.size())) { 747 argReplacements[i] = 748 fusedEntryBlock->addArgument(originalOpRegion.getArgument(i).getType()); 749 } 750 rewriter.mergeBlocks(fusedEntryBlock->getNextNode(), fusedEntryBlock, 751 argReplacements); 752 } 753 754 /// Update the body of an expanded linalg operation having index semantics. The 755 /// indices of the original operation need to be recovered by linearizing the 756 /// indices of the correspoding dimensions of the expanded operation. For now it 757 /// is assumed that the shapes of the expanded operation needed for 758 /// linearization are static. 759 static void updateExpandedIndexOpRegion(PatternRewriter &rewriter, Location loc, 760 Region &fusedRegion, 761 const ExpansionInfo &expansionInfo) { 762 // Replace the original indices by the linearization of the expanded indices. 763 for (IndexOp indexOp : 764 llvm::make_early_inc_range(fusedRegion.front().getOps<IndexOp>())) { 765 ArrayRef<int64_t> expandedDims = 766 expansionInfo.getExpandedDims(indexOp.dim()); 767 assert(!expandedDims.empty() && "expected valid expansion info"); 768 769 // Skip index operations that are not affected by the expansion. 770 if (expandedDims.size() == 1 && 771 expandedDims.front() == (int64_t)indexOp.dim()) 772 continue; 773 774 // Linearize the expanded indices of the original index dimension. 775 OpBuilder::InsertionGuard guard(rewriter); 776 rewriter.setInsertionPointAfter(indexOp); 777 ArrayRef<int64_t> expandedDimsShape = 778 expansionInfo.getExpandedShapeOfDim(indexOp.dim()).drop_front(); 779 SmallVector<Value> expandedIndices; 780 expandedIndices.reserve(expandedDims.size() - 1); 781 llvm::transform( 782 expandedDims.drop_front(), std::back_inserter(expandedIndices), 783 [&](int64_t dim) { return rewriter.create<IndexOp>(loc, dim); }); 784 Value newIndex = rewriter.create<IndexOp>(loc, expandedDims.front()); 785 for (auto it : llvm::zip(expandedDimsShape, expandedIndices)) { 786 assert(!ShapedType::isDynamic(std::get<0>(it))); 787 AffineExpr idx, acc; 788 bindDims(rewriter.getContext(), idx, acc); 789 newIndex = rewriter.create<AffineApplyOp>( 790 indexOp.getLoc(), idx + acc * std::get<0>(it), 791 ValueRange{std::get<1>(it), newIndex}); 792 } 793 rewriter.replaceOp(indexOp, newIndex); 794 } 795 } 796 797 /// Implements the fusion of a tensor_reshape op and a generic/indexed_generic 798 /// op as explained in `isFusableWithReshapeByExpansion`. Assumes that those 799 /// conditions have been satisfied. 800 static Optional<SmallVector<Value, 1>> 801 fuseWithReshapeByExpansion(LinalgOp linalgOp, TensorReshapeOp reshapeOp, 802 unsigned fusedTensorIndex, 803 PatternRewriter &rewriter) { 804 assert(isFusableWithReshapeByDimExpansion(linalgOp, fusedTensorIndex) && 805 "preconditions for fuse operation failed"); 806 // Check if reshape is expanding or collapsing. 807 bool isExpanding = 808 reshapeOp.getSrcType().getRank() < reshapeOp.getResultType().getRank(); 809 RankedTensorType expandedType = 810 isExpanding ? reshapeOp.getResultType() : reshapeOp.getSrcType(); 811 bool hasIndexSemantics = linalgOp.hasIndexSemantics() || 812 isa<IndexedGenericOp>(linalgOp.getOperation()); 813 814 ExpansionInfo expansionInfo; 815 if (failed(expansionInfo.compute(linalgOp, fusedTensorIndex, 816 reshapeOp.getReassociationMaps(), 817 expandedType.getShape()))) 818 return llvm::None; 819 820 if (hasIndexSemantics && 821 failed(isIndexedOpExpandable(linalgOp, expansionInfo))) 822 return llvm::None; 823 824 SmallVector<AffineMap, 4> expandedOpIndexingMaps = llvm::to_vector<4>( 825 llvm::map_range(linalgOp.getIndexingMaps(), [&](AffineMap m) { 826 return getIndexingMapInExpandedOp(rewriter, m, expansionInfo); 827 })); 828 829 SmallVector<Value, 4> expandedOpOperands; 830 for (auto operand : llvm::enumerate(linalgOp.getInputs())) { 831 if (operand.index() == fusedTensorIndex) { 832 expandedOpOperands.push_back(reshapeOp.src()); 833 continue; 834 } 835 AffineMap indexingMap = linalgOp.getInputIndexingMap(operand.index()); 836 RankedTensorType expandedOperandType = 837 getExpandedType(operand.value().getType().cast<RankedTensorType>(), 838 indexingMap, expansionInfo); 839 if (expandedOperandType != operand.value().getType()) { 840 // Reshape the operand to get the right type. 841 SmallVector<ReassociationIndices, 4> reassociation = 842 getReassociationForExpansion(indexingMap, expansionInfo); 843 expandedOpOperands.push_back(rewriter.create<TensorReshapeOp>( 844 linalgOp.getLoc(), expandedOperandType, operand.value(), 845 reassociation)); 846 continue; 847 } 848 expandedOpOperands.push_back(operand.value()); 849 } 850 851 Location loc = linalgOp.getLoc(); 852 SmallVector<Value, 1> outputs; 853 for (auto result : llvm::enumerate(linalgOp.getOutputs())) { 854 AffineMap indexingMap = linalgOp.getOutputIndexingMap(result.index()); 855 RankedTensorType expandedOutputType = 856 getExpandedType(result.value().getType().cast<RankedTensorType>(), 857 indexingMap, expansionInfo); 858 if (expandedOutputType != result.value().getType()) { 859 SmallVector<ReassociationIndices, 4> reassociation = 860 getReassociationForExpansion(indexingMap, expansionInfo); 861 outputs.push_back(rewriter.create<TensorReshapeOp>( 862 linalgOp.getLoc(), expandedOutputType, result.value(), 863 reassociation)); 864 } 865 } 866 867 // The iterator types of the expanded op are all parallel. 868 SmallVector<StringRef, 4> iteratorTypes(expansionInfo.getExpandedOpNumDims(), 869 getParallelIteratorTypeName()); 870 871 TypeRange resultTypes = ValueRange(outputs).getTypes(); 872 LinalgOp fusedOp = createLinalgOpOfSameType( 873 linalgOp, rewriter, linalgOp.getLoc(), resultTypes, 874 /*inputs=*/expandedOpOperands, outputs, expandedOpIndexingMaps, 875 iteratorTypes); 876 Region &fusedRegion = fusedOp->getRegion(0); 877 Region &originalRegion = linalgOp->getRegion(0); 878 879 if (isa<GenericOp>(linalgOp.getOperation())) { 880 rewriter.cloneRegionBefore(originalRegion, fusedRegion, 881 fusedRegion.begin()); 882 } else { 883 assert(isa<IndexedGenericOp>(linalgOp.getOperation())); 884 buildExpandedIndexedGenericOpRegion(rewriter, loc, originalRegion, 885 fusedRegion, expansionInfo); 886 } 887 888 // Update the index accesses after the expansion. 889 if (linalgOp.hasIndexSemantics()) 890 updateExpandedIndexOpRegion(rewriter, loc, fusedRegion, expansionInfo); 891 892 // Reshape the result values to their original shape if this is a collapsing 893 // reshape folded into its consumer. 894 SmallVector<Value, 1> resultVals; 895 for (auto result : llvm::enumerate(linalgOp->getResults())) { 896 if (!isExpanding && 897 resultTypes[result.index()] != result.value().getType()) { 898 SmallVector<ReassociationIndices, 4> reassociation = 899 getReassociationForExpansion( 900 linalgOp.getOutputIndexingMap(result.index()), expansionInfo); 901 resultVals.push_back(rewriter.create<TensorReshapeOp>( 902 linalgOp.getLoc(), result.value().getType(), 903 fusedOp->getResult(result.index()), reassociation)); 904 } else { 905 resultVals.push_back(fusedOp->getResult(result.index())); 906 } 907 } 908 // Assuming a single result. 909 return resultVals; 910 } 911 912 namespace { 913 914 /// Pattern to fold tensor_reshape op with its consumer by using the source of 915 /// the reshape op as the operand in the consumer (instead of the result of the 916 /// tensor_reshapeop) when the tensor_reshape op is collapsing. The 917 /// corresponding index map in the consumer needs to be modified to linearize 918 /// the folded dimension. 919 /// 920 /// For example, 921 /// 922 /// #map0 = affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)> 923 /// %0 = linalg.tensor_reshape %arg0 924 /// [affine_map<(i, j, k, l) -> (i)>, affine_map<(i, j, k, l) -> (j, k)>, 925 /// affine_map<(i, j, k, l) -> (l)>] 926 /// tensor<?x?x?xf32> into tensor<?x?x4x?xf32> 927 /// %1 = linalg.generic { indexing_maps = [#map0, #map0, #map0], ... } 928 /// ins(%0, %arg1 : tensor<?x?x4x?xf32>, tensor<?x?x4x?xf32>) ... 929 /// -> tensor<?x?x4x?xf32> 930 /// 931 /// can be folded into 932 /// 933 /// #map0 = affine_map<(d0, d1, d2, d3) -> (d0, d1 * 4 + d2, d3)> 934 /// #map1 = affine_map<(d0, d1, d2, d3) -> (d0, d1, d2, d3)> 935 /// %0 = linalg.generic { indexing_maps = [#map0, #map1, #map1] ... } 936 /// ins(%arg0, %arg1 : tensor<?x?x?xf32>, tensor<?x?x4x?xf32>) ... 937 /// -> tensor<?x?x4x?xf32> 938 template <typename LinalgOpTy, bool foldUnitDimReshapesOnly> 939 struct FoldProducerReshapeOpByLinearization 940 : public OpRewritePattern<LinalgOpTy> { 941 using OpRewritePattern<LinalgOpTy>::OpRewritePattern; 942 943 LogicalResult matchAndRewrite(LinalgOpTy op, 944 PatternRewriter &rewriter) const override { 945 if (!op.hasTensorSemantics()) 946 return failure(); 947 LinalgOp linalgOp = cast<LinalgOp>(op.getOperation()); 948 for (auto operand : llvm::enumerate(linalgOp.getInputs())) { 949 TensorReshapeOp reshapeOp = 950 operand.value().getDefiningOp<TensorReshapeOp>(); 951 if (!reshapeOp || 952 !isTensorReshapeOpFoldableByLinearization( 953 reshapeOp, linalgOp.getInputIndexingMap(operand.index()), 954 /*asProducer =*/true) || 955 (foldUnitDimReshapesOnly && 956 !isUnitDimExpansionOnly(reshapeOp.getResultType().getShape(), 957 reshapeOp.getReassociationMaps()))) 958 continue; 959 960 // Compute the fused operands list, 961 SmallVector<Value, 2> fusedOperands(linalgOp.getInputs()); 962 fusedOperands[operand.index()] = reshapeOp.src(); 963 fusedOperands.append(linalgOp.getOutputs().begin(), 964 linalgOp.getOutputs().end()); 965 966 // Compute indexing_maps for the fused operation. The indexing_maps for 967 // the operands of the consumers that arent fused are the same. 968 SmallVector<AffineMap, 4> fusedIndexMaps = llvm::to_vector<4>( 969 op.indexing_maps().template getAsValueRange<AffineMapAttr>()); 970 971 // Accepted consumer maps are either identity or permutation. 972 auto invMap = inversePermutation(fusedIndexMaps[operand.index()]); 973 974 // Compute the indexing map to use for the result of the producer. 975 AffineMap modifiedMap = 976 linearizeCollapsedDims(invMap, reshapeOp.getResultType().getShape(), 977 reshapeOp.getReassociationMaps()); 978 for (AffineExpr expr : modifiedMap.getResults()) { 979 if (!expr.isPureAffine()) 980 return failure(); 981 } 982 fusedIndexMaps[operand.index()] = modifiedMap; 983 984 // Further check that the resulting index maps can be fused and 985 // inverted. Without this the resultant op is not legal. 986 if (!inversePermutation(concatAffineMaps(fusedIndexMaps))) { 987 return rewriter.notifyMatchFailure( 988 op, "fused op loop bound computation failed"); 989 } 990 991 rewriter.startRootUpdate(op); 992 op->setOperands(fusedOperands); 993 op.indexing_mapsAttr(rewriter.getAffineMapArrayAttr(fusedIndexMaps)); 994 rewriter.finalizeRootUpdate(op); 995 return success(); 996 } 997 return failure(); 998 } 999 }; 1000 1001 static SmallVector<ReassociationIndices> 1002 getReassociationIndices(ArrayRef<AffineMap> maps) { 1003 SmallVector<ReassociationIndices> reassociation; 1004 for (AffineMap map : maps) { 1005 ReassociationIndices indices; 1006 for (unsigned i = 0, e = map.getNumResults(); i < e; i++) { 1007 unsigned pos = map.getResult(i).cast<AffineDimExpr>().getPosition(); 1008 indices.push_back(pos); 1009 } 1010 reassociation.push_back(indices); 1011 } 1012 return reassociation; 1013 } 1014 1015 /// Pattern to move rank reducing reshape after an elementwise linalg generic 1016 /// op. This is useful to expose more fusion opportunities between named ops and 1017 /// generic op. This can only be done if there is no broadcast or permuation 1018 /// within the dimensions we need to merge. 1019 /// 1020 /// For example, 1021 /// 1022 /// %0 = linalg.tensor_reshape %A [ 1023 /// affine_map<(d0, d1, d2) -> (d0, d1)>, affine_map<(d0, d1, d2) -> (d2)>] 1024 /// : tensor<12544x16xf32> into tensor<112x112x16xf32> 1025 /// %2 = linalg.generic {indexing_maps = [ 1026 /// affine_map<(d0, d1, d2) -> (d0, d1, d2)>, 1027 /// affine_map<(d0, d1, d2) -> (d2)>, 1028 /// affine_map<(d0, d1, d2) -> (d0, d1, d2)>], iterator_types = 1029 /// ["parallel", "parallel", "parallel"]} { 1030 /// } -> tensor<112x112x16xf32> 1031 /// 1032 /// into 1033 /// 1034 /// %2 = linalg.generic {indexing_maps = [ 1035 /// affine_map<(d0, d1) -> (d0, d1)>, 1036 /// affine_map<(d0, d1) -> (d1)>, 1037 /// affine_map<(d0, d1) -> (d0, d1)>], 1038 /// iterator_types = ["parallel", "parallel"]} ins(%arg0, %arg1 1039 /// : tensor<12544x16xf32>, tensor<16xf32>) outs(%1 : tensor<12544x16xf32>) { 1040 /// } -> tensor<12544x16xf32> 1041 /// %3 = linalg.tensor_reshape %2 [ 1042 /// #affine_map<(d0, d1, d2) -> (d0, d1)>, affine_map<(d0, d1, d2) -> (d2)>] 1043 /// : tensor<12544x16xf32> into tensor<112x112x16xf32> 1044 template <typename GenericOpTy> 1045 struct PushExpandingReshape : public OpRewritePattern<GenericOpTy> { 1046 using OpRewritePattern<GenericOpTy>::OpRewritePattern; 1047 1048 LogicalResult matchAndRewrite(GenericOpTy op, 1049 PatternRewriter &rewriter) const override { 1050 // Only apply to elementwise linalg on tensor. 1051 if (!op.hasTensorSemantics() || 1052 op.getNumParallelLoops() != op.getNumLoops()) 1053 return failure(); 1054 // Only support identity output maps. It could be extended to permuations if 1055 // needed. 1056 if (llvm::any_of(op.getOutputIndexingMaps(), 1057 [](AffineMap map) { return !map.isIdentity(); })) 1058 return failure(); 1059 int64_t destRank = op.getNumParallelLoops(); 1060 SmallVector<Value, 4> newOperands = llvm::to_vector<4>(op.getInputs()); 1061 TensorReshapeOp reshapeFound; 1062 // 1. Look for tensor_reshape operands and figure out save the dimensions 1063 // merged. 1064 for (auto operand : llvm::enumerate(op.getInputs())) { 1065 TensorReshapeOp reshapeOp = 1066 operand.value().template getDefiningOp<TensorReshapeOp>(); 1067 if (!reshapeOp || reshapeOp.getSrcType().getRank() > 1068 reshapeOp.getResultType().getRank()) { 1069 continue; 1070 } 1071 // TODO: We could support non-identity map as long as the merged 1072 // dimensions are still contiguous. 1073 if (!op.getIndexingMaps()[operand.index()].isIdentity()) 1074 continue; 1075 if (reshapeFound) { 1076 // Only support a second reshape op if it has the same reassociate maps. 1077 if (reshapeFound.getReassociationMaps() == 1078 reshapeOp.getReassociationMaps()) 1079 newOperands[operand.index()] = reshapeOp.src(); 1080 continue; 1081 } 1082 reshapeFound = reshapeOp; 1083 newOperands[operand.index()] = reshapeOp.src(); 1084 } 1085 if (!reshapeFound) 1086 return failure(); 1087 1088 // Calculate the reassociation indices and rassociated reverse map. 1089 SmallVector<ReassociationIndices> reassociation = 1090 getReassociationIndices(reshapeFound.getReassociationMaps()); 1091 SmallVector<unsigned, 4> remap(destRank); 1092 for (auto &indices : llvm::enumerate(reassociation)) { 1093 for (int64_t index : indices.value()) { 1094 remap[index] = indices.index(); 1095 } 1096 } 1097 // 2. Verify that we can merge the dimensions in the linalg and that we 1098 // don't need to create new reshapes operands. Inserting new reshape 1099 // operands would defeat the purpose of the transformation. 1100 for (auto operand : llvm::enumerate(op.getInputs())) { 1101 if (operand.value() == newOperands[operand.index()]) { 1102 AffineMap map = op.getIndexingMaps()[operand.index()]; 1103 for (unsigned i : llvm::seq(unsigned(0), map.getNumResults())) { 1104 if (reassociation[remap[map.getDimPosition(i)]].size() > 1) 1105 return failure(); 1106 } 1107 } 1108 } 1109 1110 // 3. Calculate the affine map remapping and the reassociation to apply to 1111 // output tensors. 1112 SmallVector<AffineMap, 4> newMaps; 1113 unsigned newRank = reassociation.size(); 1114 for (auto map : op.getIndexingMaps()) { 1115 SmallVector<AffineExpr> newExprs; 1116 for (auto expr : map.getResults()) { 1117 unsigned position = expr.template cast<AffineDimExpr>().getPosition(); 1118 // Skip dimension merged except for the last of the group. 1119 if (reassociation[remap[position]].back() == position) { 1120 newExprs.push_back( 1121 getAffineDimExpr(remap[position], op.getContext())); 1122 } 1123 } 1124 newMaps.push_back(AffineMap::get(newRank, 0, newExprs, op.getContext())); 1125 } 1126 1127 // 4. Reshape the output tensors. 1128 SmallVector<Value> newOutputs; 1129 SmallVector<Type> newOutputTypes; 1130 for (auto output : op.outputs()) { 1131 Value newOutput = rewriter.create<TensorReshapeOp>( 1132 op->getLoc(), reshapeFound.getSrcType(), output, reassociation); 1133 newOutputTypes.push_back(newOutput.getType()); 1134 newOutputs.push_back(newOutput); 1135 } 1136 // 5. Create a new generic op with lowerer rank. 1137 SmallVector<StringRef, 4> iteratorTypes(newRank, 1138 getParallelIteratorTypeName()); 1139 auto newOp = 1140 rewriter.create<GenericOpTy>(op->getLoc(), newOutputTypes, newOperands, 1141 newOutputs, newMaps, iteratorTypes); 1142 rewriter.inlineRegionBefore(op.region(), newOp.region(), 1143 newOp.region().begin()); 1144 // 6. Reshape the so that the type matches the uses. 1145 SmallVector<Value> newResults; 1146 for (auto result : llvm::enumerate(newOp->getResults())) { 1147 newResults.push_back(rewriter.create<TensorReshapeOp>( 1148 op->getLoc(), op.getOutputTensorTypes()[result.index()], 1149 result.value(), reassociation)); 1150 } 1151 rewriter.replaceOp(op, newResults); 1152 return success(); 1153 } 1154 }; 1155 1156 /// Pattern to fuse a tensor_reshape op with its consumer 1157 /// generic/indexed_generic op, when the reshape op is collapsing 1158 /// dimensions. The dimensionality of the loop in the consumer is expanded. 1159 template <typename GenericOpTy> 1160 class FoldWithProducerReshapeOpByExpansion 1161 : public OpRewritePattern<GenericOpTy> { 1162 public: 1163 FoldWithProducerReshapeOpByExpansion(MLIRContext *context, 1164 bool foldUnitDimReshapes, 1165 PatternBenefit benefit = 1) 1166 : OpRewritePattern<GenericOpTy>(context, benefit), 1167 allowFoldingUnitDimReshapes(foldUnitDimReshapes) {} 1168 1169 LogicalResult matchAndRewrite(GenericOpTy genericOp, 1170 PatternRewriter &rewriter) const override { 1171 LinalgOp linalgOp = cast<LinalgOp>(genericOp.getOperation()); 1172 for (auto operand : llvm::enumerate(linalgOp.getInputs())) { 1173 TensorReshapeOp reshapeOp = 1174 operand.value().getDefiningOp<TensorReshapeOp>(); 1175 if (!reshapeOp) 1176 continue; 1177 1178 // Fold only if 1179 // - The tensor reshape op is folding. 1180 // - All constraints of fusing with reshape by expansion are met. 1181 if (reshapeOp.getSrcType().getRank() < 1182 reshapeOp.getResultType().getRank() || 1183 !isFusableWithReshapeByDimExpansion(linalgOp, operand.index()) || 1184 (!allowFoldingUnitDimReshapes && 1185 isUnitDimExpansionOnly(reshapeOp.getSrcType().getShape(), 1186 reshapeOp.getReassociationMaps()))) 1187 continue; 1188 1189 Optional<SmallVector<Value, 1>> replacementValues = 1190 fuseWithReshapeByExpansion(linalgOp, reshapeOp, operand.index(), 1191 rewriter); 1192 if (!replacementValues) 1193 return failure(); 1194 rewriter.replaceOp(genericOp, replacementValues.getValue()); 1195 return success(); 1196 } 1197 return failure(); 1198 } 1199 1200 private: 1201 bool allowFoldingUnitDimReshapes; 1202 }; 1203 1204 /// Pattern to fold tensor_reshape op with its producer. The corresponding index 1205 /// map in the consumer needs to be modified to linearize the folded dimension. 1206 template <bool foldUnitDimReshapesOnly> 1207 struct FoldConsumerReshapeOpByLinearization 1208 : public OpRewritePattern<TensorReshapeOp> { 1209 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 1210 1211 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 1212 PatternRewriter &rewriter) const override { 1213 LinalgOp producer = reshapeOp.src().getDefiningOp<LinalgOp>(); 1214 if (!producer || 1215 !isa<GenericOp, IndexedGenericOp>(producer.getOperation()) || 1216 !producer.hasTensorSemantics() || producer.getNumOutputs() != 1 || 1217 !isTensorReshapeOpFoldableByLinearization( 1218 reshapeOp, producer.getOutputIndexingMap(0), 1219 /*asProducer =*/false) || 1220 (foldUnitDimReshapesOnly && 1221 !isUnitDimExpansionOnly(reshapeOp.getSrcType().getShape(), 1222 reshapeOp.getReassociationMaps()))) 1223 return failure(); 1224 // The indexing_maps for the operands of the fused operation are same as 1225 // those for the operands of the producer. 1226 SmallVector<AffineMap, 4> fusedIndexMaps = llvm::to_vector<4>( 1227 producer.indexing_maps().getAsValueRange<AffineMapAttr>()); 1228 1229 auto invMap = inversePermutation(producer.getOutputIndexingMap(0)); 1230 1231 // Compute the indexing map to use for the operand of the producer. 1232 AffineMap modifiedMap = 1233 linearizeCollapsedDims(invMap, reshapeOp.getSrcType().getShape(), 1234 reshapeOp.getReassociationMaps()); 1235 for (AffineExpr expr : modifiedMap.getResults()) { 1236 if (!expr.isPureAffine()) { 1237 return rewriter.notifyMatchFailure( 1238 producer, "fused op indexing map is not affine"); 1239 } 1240 } 1241 fusedIndexMaps.back() = modifiedMap; 1242 1243 // Further check that the resulting index maps can be fused and 1244 // inverted. Without this the resultant op is not legal. 1245 if (!inversePermutation(concatAffineMaps(fusedIndexMaps))) { 1246 return rewriter.notifyMatchFailure( 1247 producer, "fused op loop bound computation failed"); 1248 } 1249 1250 Location loc = producer.getLoc(); 1251 Value output = rewriter.create<TensorReshapeOp>( 1252 loc, producer.getOutputs()[0], reshapeOp.getReassociationExprs()); 1253 LinalgOp fusedOp = createLinalgOpOfSameType( 1254 producer, rewriter, loc, reshapeOp.getResultType(), 1255 /*inputs=*/producer.getInputs(), 1256 // TODO: handle outputs. 1257 /*outputs=*/output, rewriter.getAffineMapArrayAttr(fusedIndexMaps), 1258 producer.iterator_types(), 1259 /*doc=*/nullptr, 1260 /*library_call=*/nullptr, 1261 /*sparse=*/nullptr); 1262 auto &fusedRegion = fusedOp->getRegion(0); 1263 rewriter.cloneRegionBefore(producer->getRegion(0), fusedRegion, 1264 fusedRegion.begin()); 1265 rewriter.replaceOp(reshapeOp, fusedOp->getResults()); 1266 return success(); 1267 } 1268 }; 1269 1270 /// Pattern to fold a tensor_reshape op with its producer generic op if the 1271 /// tensor_reshape op is expanding, by expanding the dimensionality of the loop 1272 /// in the producer op. 1273 struct FoldReshapeWithGenericOpByExpansion 1274 : public OpRewritePattern<TensorReshapeOp> { 1275 using OpRewritePattern<TensorReshapeOp>::OpRewritePattern; 1276 LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp, 1277 PatternRewriter &rewriter) const override { 1278 // Fold only if 1279 // - The tensor reshape op is a expanding case. 1280 // - All constraints of fusing with reshape by expansion are met. 1281 if (reshapeOp.getSrcType().getRank() > reshapeOp.getResultType().getRank()) 1282 return failure(); 1283 LinalgOp producer = reshapeOp.src().getDefiningOp<LinalgOp>(); 1284 if (!producer || producer.getNumOutputs() != 1 || 1285 !isFusableWithReshapeByDimExpansion(producer, 1286 producer.getNumInputs()) || 1287 isUnitDimExpansionOnly(reshapeOp.getResultType().getShape(), 1288 reshapeOp.getReassociationMaps())) 1289 return failure(); 1290 Optional<SmallVector<Value, 1>> replacementValues = 1291 fuseWithReshapeByExpansion(producer, reshapeOp, producer.getNumInputs(), 1292 rewriter); 1293 if (!replacementValues) 1294 return failure(); 1295 rewriter.replaceOp(reshapeOp, replacementValues.getValue()); 1296 return success(); 1297 } 1298 }; 1299 1300 /// Pattern to fold a GenericOp/IndexedGenericOp with a splat constant. 1301 template <typename LinalgOpTy> 1302 class FoldSplatConstants : public OpRewritePattern<LinalgOpTy> { 1303 public: 1304 FoldSplatConstants(MLIRContext *context, ControlElementwiseOpsFusionFn &fun, 1305 PatternBenefit benefit = 1) 1306 : OpRewritePattern<LinalgOpTy>(context, benefit), controlFn(fun) {} 1307 1308 LogicalResult matchAndRewrite(LinalgOpTy op, 1309 PatternRewriter &rewriter) const override { 1310 if (!op.hasTensorSemantics()) 1311 return failure(); 1312 LinalgOp linalgOp = cast<LinalgOp>(op.getOperation()); 1313 for (auto operand : llvm::enumerate(linalgOp.getInputOpOperands())) { 1314 ConstantOp constantOp = operand.value().get().getDefiningOp<ConstantOp>(); 1315 if (!constantOp || 1316 !constantOp.value().cast<DenseElementsAttr>().isSplat() || 1317 !controlFn(constantOp->getResult(0), operand.value())) 1318 continue; 1319 1320 // The indexing_maps for the operands of the fused operation are same as 1321 // those for the operands of the linalgOp without the indexing map at 1322 // operand.index() 1323 SmallVector<AffineMap, 4> fusedIndexMaps = llvm::to_vector<4>( 1324 linalgOp.indexing_maps().getAsValueRange<AffineMapAttr>()); 1325 fusedIndexMaps.erase(std::next(fusedIndexMaps.begin(), operand.index())); 1326 1327 // Check if the operation shapes to loops map is computable. 1328 if (!inversePermutation(concatAffineMaps(fusedIndexMaps))) { 1329 return rewriter.notifyMatchFailure( 1330 linalgOp, "fused op loop bound computation failed"); 1331 } 1332 1333 // The operands list is same as the linalgOp with the argument for 1334 // constant index dropped. 1335 SmallVector<Value, 4> fusedOperands(linalgOp.getInputs()); 1336 fusedOperands.erase(std::next(fusedOperands.begin(), operand.index())); 1337 1338 // Create a constant scalar value from the splat constant. 1339 Value scalarConstant = rewriter.create<ConstantOp>( 1340 constantOp.getLoc(), 1341 constantOp.value().cast<DenseElementsAttr>().getSplatValue()); 1342 1343 LinalgOp fusedOp = createLinalgOpOfSameType( 1344 linalgOp, rewriter, rewriter.getUnknownLoc(), 1345 linalgOp->getResultTypes(), 1346 /*inputs=*/fusedOperands, 1347 /*outputs=*/linalgOp.getOutputs(), 1348 rewriter.getAffineMapArrayAttr(fusedIndexMaps), 1349 linalgOp.iterator_types(), 1350 /*doc=*/nullptr, 1351 /*library_call=*/nullptr, 1352 /*sparse=*/nullptr); 1353 1354 // Map the block argument corresponding to the replaced argument with the 1355 // scalar constant. 1356 Region &linalgOpRegion = linalgOp->getRegion(0); 1357 Block &entryBlock = *linalgOpRegion.begin(); 1358 unsigned argIndex = entryBlock.getNumArguments() - 1359 linalgOp.getNumShapedOperands() + operand.index(); 1360 BlockAndValueMapping mapping; 1361 mapping.map(entryBlock.getArgument(argIndex), scalarConstant); 1362 Region &fusedRegion = fusedOp->getRegion(0); 1363 rewriter.cloneRegionBefore(linalgOpRegion, fusedRegion, 1364 fusedRegion.begin(), mapping); 1365 rewriter.replaceOp(linalgOp, fusedOp->getResults()); 1366 return success(); 1367 } 1368 return failure(); 1369 } 1370 1371 private: 1372 ControlElementwiseOpsFusionFn controlFn; 1373 }; 1374 } // namespace 1375 1376 static Optional<SmallVector<Value, 1>> 1377 fuseElementwiseOps(PatternRewriter &rewriter, OpOperand &consumerOpOperand, 1378 const ControlElementwiseOpsFusionFn &controlFn) { 1379 Operation *producer = consumerOpOperand.get().getDefiningOp(); 1380 if (!producer || producer->getNumResults() != 1) 1381 return llvm::None; 1382 1383 // Fuse when consumer is GenericOp or IndexedGenericOp. 1384 if (!isa<GenericOp, IndexedGenericOp>(consumerOpOperand.getOwner()) || 1385 !isa<GenericOp, IndexedGenericOp>(producer)) 1386 return llvm::None; 1387 1388 return fuseElementwiseOpsImpl(cast<LinalgOp>(producer), consumerOpOperand, 1389 controlFn, rewriter); 1390 } 1391 1392 namespace { 1393 /// Patterns to fuse a generic op, with the producer of its operands. 1394 template <typename LinalgOpTy> 1395 class FuseElementwiseOps : public OpRewritePattern<LinalgOpTy> { 1396 public: 1397 FuseElementwiseOps(MLIRContext *context, ControlElementwiseOpsFusionFn &fun, 1398 PatternBenefit benefit = 1) 1399 : OpRewritePattern<LinalgOpTy>(context, benefit), controlFn(fun) {} 1400 1401 LogicalResult matchAndRewrite(LinalgOpTy op, 1402 PatternRewriter &rewriter) const override { 1403 // Find the first operand that is defined by another generic op on tensors. 1404 for (OpOperand &opOperand : op.getShapedOpOperands()) { 1405 LinalgOp producerOp = 1406 dyn_cast_or_null<LinalgOp>(opOperand.get().getDefiningOp()); 1407 if (!producerOp || !producerOp.hasTensorSemantics()) 1408 continue; 1409 Optional<SmallVector<Value, 1>> fusedOpResults = 1410 fuseElementwiseOps(rewriter, opOperand, controlFn); 1411 if (fusedOpResults) { 1412 rewriter.replaceOp(op, *fusedOpResults); 1413 return success(); 1414 } 1415 } 1416 return failure(); 1417 } 1418 1419 private: 1420 ControlElementwiseOpsFusionFn controlFn; 1421 }; 1422 1423 /// Pass that fuses generic ops on tensors. Used only for testing. 1424 struct FusionOfTensorOpsPass 1425 : public LinalgFusionOfTensorOpsBase<FusionOfTensorOpsPass> { 1426 void runOnOperation() override { 1427 Operation *op = getOperation(); 1428 RewritePatternSet patterns(op->getContext()); 1429 populateElementwiseOpsFusionPatterns( 1430 patterns, 1431 LinalgElementwiseFusionOptions().setAllowFoldingUnitDimReshapes( 1432 allowFoldingUnitDimReshapes)); 1433 (void)applyPatternsAndFoldGreedily(op->getRegions(), std::move(patterns)); 1434 } 1435 }; 1436 1437 /// Pass to test folding of reshape op with generic/indexed_generic ops by 1438 /// linearization. 1439 struct FoldReshapeOpsByLinearizationPass 1440 : public LinalgFoldReshapeOpsByLinearizationBase< 1441 FoldReshapeOpsByLinearizationPass> { 1442 void runOnOperation() override { 1443 Operation *op = getOperation(); 1444 RewritePatternSet patterns(op->getContext()); 1445 populateFoldReshapeOpsByLinearizationPatterns(patterns); 1446 (void)applyPatternsAndFoldGreedily(op->getRegions(), std::move(patterns)); 1447 } 1448 }; 1449 1450 } // namespace 1451 1452 void mlir::linalg::populateFoldReshapeOpsByLinearizationPatterns( 1453 RewritePatternSet &patterns) { 1454 patterns.add<FoldProducerReshapeOpByLinearization<GenericOp, false>, 1455 FoldProducerReshapeOpByLinearization<IndexedGenericOp, false>, 1456 FoldConsumerReshapeOpByLinearization<false>>( 1457 patterns.getContext()); 1458 } 1459 1460 void mlir::linalg::populateFoldUnitDimsReshapeOpsByLinearizationPatterns( 1461 RewritePatternSet &patterns) { 1462 patterns.add<FoldProducerReshapeOpByLinearization<GenericOp, true>, 1463 FoldProducerReshapeOpByLinearization<IndexedGenericOp, true>, 1464 FoldConsumerReshapeOpByLinearization<true>>( 1465 patterns.getContext()); 1466 } 1467 1468 void mlir::linalg::populateFoldReshapeOpsByExpansionPatterns( 1469 RewritePatternSet &patterns, bool allowFoldingUnitDimReshapes) { 1470 patterns.add<FoldReshapeWithGenericOpByExpansion>(patterns.getContext()); 1471 patterns.add<FoldWithProducerReshapeOpByExpansion<GenericOp>, 1472 FoldWithProducerReshapeOpByExpansion<IndexedGenericOp>>( 1473 patterns.getContext(), allowFoldingUnitDimReshapes); 1474 } 1475 1476 void mlir::linalg::populateElementwiseOpsFusionPatterns( 1477 RewritePatternSet &patterns, LinalgElementwiseFusionOptions options) { 1478 auto *context = patterns.getContext(); 1479 patterns 1480 .add<FuseElementwiseOps<GenericOp>, FuseElementwiseOps<IndexedGenericOp>, 1481 FoldSplatConstants<GenericOp>, FoldSplatConstants<IndexedGenericOp>>( 1482 context, options.controlElementwiseOpsFusionFn); 1483 populateFoldReshapeOpsByExpansionPatterns( 1484 patterns, options.allowFoldingUnitDimReshapes); 1485 AffineApplyOp::getCanonicalizationPatterns(patterns, context); 1486 GenericOp::getCanonicalizationPatterns(patterns, context); 1487 IndexedGenericOp::getCanonicalizationPatterns(patterns, context); 1488 TensorReshapeOp::getCanonicalizationPatterns(patterns, context); 1489 } 1490 1491 void mlir::linalg::populatePushReshapeOpsPatterns(RewritePatternSet &patterns) { 1492 auto *context = patterns.getContext(); 1493 patterns.add<PushExpandingReshape<GenericOp>, 1494 PushExpandingReshape<IndexedGenericOp>>(context); 1495 } 1496 1497 std::unique_ptr<Pass> mlir::createLinalgFusionOfTensorOpsPass() { 1498 return std::make_unique<FusionOfTensorOpsPass>(); 1499 } 1500 1501 std::unique_ptr<Pass> mlir::createFoldReshapeOpsByLinearizationPass() { 1502 return std::make_unique<FoldReshapeOpsByLinearizationPass>(); 1503 } 1504