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 /// Pattern to fuse a tensor_reshape op with its consumer
1002 /// generic/indexed_generic op, when the reshape op is collapsing
1003 /// dimensions. The dimensionality of the loop in the consumer is expanded.
1004 template <typename GenericOpTy>
1005 class FoldWithProducerReshapeOpByExpansion
1006     : public OpRewritePattern<GenericOpTy> {
1007 public:
1008   FoldWithProducerReshapeOpByExpansion(MLIRContext *context,
1009                                        bool foldUnitDimReshapes,
1010                                        PatternBenefit benefit = 1)
1011       : OpRewritePattern<GenericOpTy>(context, benefit),
1012         allowFoldingUnitDimReshapes(foldUnitDimReshapes) {}
1013 
1014   LogicalResult matchAndRewrite(GenericOpTy genericOp,
1015                                 PatternRewriter &rewriter) const override {
1016     LinalgOp linalgOp = cast<LinalgOp>(genericOp.getOperation());
1017     for (auto operand : llvm::enumerate(linalgOp.getInputs())) {
1018       TensorReshapeOp reshapeOp =
1019           operand.value().getDefiningOp<TensorReshapeOp>();
1020       if (!reshapeOp)
1021         continue;
1022 
1023       // Fold only if
1024       // - The tensor reshape op is folding.
1025       // - All constraints of fusing with reshape by expansion are met.
1026       if (reshapeOp.getSrcType().getRank() <
1027               reshapeOp.getResultType().getRank() ||
1028           !isFusableWithReshapeByDimExpansion(linalgOp, operand.index()) ||
1029           (!allowFoldingUnitDimReshapes &&
1030            isUnitDimExpansionOnly(reshapeOp.getSrcType().getShape(),
1031                                   reshapeOp.getReassociationMaps())))
1032         continue;
1033 
1034       Optional<SmallVector<Value, 1>> replacementValues =
1035           fuseWithReshapeByExpansion(linalgOp, reshapeOp, operand.index(),
1036                                      rewriter);
1037       if (!replacementValues)
1038         return failure();
1039       rewriter.replaceOp(genericOp, replacementValues.getValue());
1040       return success();
1041     }
1042     return failure();
1043   }
1044 
1045 private:
1046   bool allowFoldingUnitDimReshapes;
1047 };
1048 
1049 /// Pattern to fold tensor_reshape op with its producer. The corresponding index
1050 /// map in the consumer needs to be modified to linearize the folded dimension.
1051 template <bool foldUnitDimReshapesOnly>
1052 struct FoldConsumerReshapeOpByLinearization
1053     : public OpRewritePattern<TensorReshapeOp> {
1054   using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
1055 
1056   LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
1057                                 PatternRewriter &rewriter) const override {
1058     LinalgOp producer = reshapeOp.src().getDefiningOp<LinalgOp>();
1059     if (!producer ||
1060         !isa<GenericOp, IndexedGenericOp>(producer.getOperation()) ||
1061         !producer.hasTensorSemantics() || producer.getNumOutputs() != 1 ||
1062         !isTensorReshapeOpFoldableByLinearization(
1063             reshapeOp, producer.getOutputIndexingMap(0),
1064             /*asProducer =*/false) ||
1065         (foldUnitDimReshapesOnly &&
1066          !isUnitDimExpansionOnly(reshapeOp.getSrcType().getShape(),
1067                                  reshapeOp.getReassociationMaps())))
1068       return failure();
1069     // The indexing_maps for the operands of the fused operation are same as
1070     // those for the operands of the producer.
1071     SmallVector<AffineMap, 4> fusedIndexMaps = llvm::to_vector<4>(
1072         producer.indexing_maps().getAsValueRange<AffineMapAttr>());
1073 
1074     auto invMap = inversePermutation(producer.getOutputIndexingMap(0));
1075 
1076     // Compute the indexing map to use for the operand of the producer.
1077     AffineMap modifiedMap =
1078         linearizeCollapsedDims(invMap, reshapeOp.getSrcType().getShape(),
1079                                reshapeOp.getReassociationMaps());
1080     for (AffineExpr expr : modifiedMap.getResults()) {
1081       if (!expr.isPureAffine()) {
1082         return rewriter.notifyMatchFailure(
1083             producer, "fused op indexing map is not affine");
1084       }
1085     }
1086     fusedIndexMaps.back() = modifiedMap;
1087 
1088     // Further check that the resulting index maps can be fused and
1089     // inverted. Without this the resultant op is not legal.
1090     if (!inversePermutation(concatAffineMaps(fusedIndexMaps))) {
1091       return rewriter.notifyMatchFailure(
1092           producer, "fused op loop bound computation failed");
1093     }
1094 
1095     Location loc = producer.getLoc();
1096     Value output = rewriter.create<TensorReshapeOp>(
1097         loc, producer.getOutputs()[0], reshapeOp.getReassociationExprs());
1098     LinalgOp fusedOp = createLinalgOpOfSameType(
1099         producer, rewriter, loc, reshapeOp.getResultType(),
1100         /*inputs=*/producer.getInputs(),
1101         // TODO: handle outputs.
1102         /*outputs=*/output, rewriter.getAffineMapArrayAttr(fusedIndexMaps),
1103         producer.iterator_types(),
1104         /*doc=*/nullptr,
1105         /*library_call=*/nullptr,
1106         /*sparse=*/nullptr);
1107     auto &fusedRegion = fusedOp->getRegion(0);
1108     rewriter.cloneRegionBefore(producer->getRegion(0), fusedRegion,
1109                                fusedRegion.begin());
1110     rewriter.replaceOp(reshapeOp, fusedOp->getResults());
1111     return success();
1112   }
1113 };
1114 
1115 /// Pattern to fold a tensor_reshape op with its producer generic op if the
1116 /// tensor_reshape op is expanding, by expanding the dimensionality of the loop
1117 /// in the producer op.
1118 struct FoldReshapeWithGenericOpByExpansion
1119     : public OpRewritePattern<TensorReshapeOp> {
1120   using OpRewritePattern<TensorReshapeOp>::OpRewritePattern;
1121   LogicalResult matchAndRewrite(TensorReshapeOp reshapeOp,
1122                                 PatternRewriter &rewriter) const override {
1123     // Fold only if
1124     // - The tensor reshape op is a expanding case.
1125     // - All constraints of fusing with reshape by expansion are met.
1126     if (reshapeOp.getSrcType().getRank() > reshapeOp.getResultType().getRank())
1127       return failure();
1128     LinalgOp producer = reshapeOp.src().getDefiningOp<LinalgOp>();
1129     if (!producer || producer.getNumOutputs() != 1 ||
1130         !isFusableWithReshapeByDimExpansion(producer,
1131                                             producer.getNumInputs()) ||
1132         isUnitDimExpansionOnly(reshapeOp.getResultType().getShape(),
1133                                reshapeOp.getReassociationMaps()))
1134       return failure();
1135     Optional<SmallVector<Value, 1>> replacementValues =
1136         fuseWithReshapeByExpansion(producer, reshapeOp, producer.getNumInputs(),
1137                                    rewriter);
1138     if (!replacementValues)
1139       return failure();
1140     rewriter.replaceOp(reshapeOp, replacementValues.getValue());
1141     return success();
1142   }
1143 };
1144 
1145 /// Pattern to fold a GenericOp/IndexedGenericOp with a splat constant.
1146 template <typename LinalgOpTy>
1147 class FoldSplatConstants : public OpRewritePattern<LinalgOpTy> {
1148 public:
1149   FoldSplatConstants(MLIRContext *context, ControlElementwiseOpsFusionFn &fun,
1150                      PatternBenefit benefit = 1)
1151       : OpRewritePattern<LinalgOpTy>(context, benefit), controlFn(fun) {}
1152 
1153   LogicalResult matchAndRewrite(LinalgOpTy op,
1154                                 PatternRewriter &rewriter) const override {
1155     if (!op.hasTensorSemantics())
1156       return failure();
1157     LinalgOp linalgOp = cast<LinalgOp>(op.getOperation());
1158     for (auto operand : llvm::enumerate(linalgOp.getInputOpOperands())) {
1159       ConstantOp constantOp = operand.value().get().getDefiningOp<ConstantOp>();
1160       if (!constantOp ||
1161           !constantOp.value().cast<DenseElementsAttr>().isSplat() ||
1162           !controlFn(constantOp->getResult(0), operand.value()))
1163         continue;
1164 
1165       // The indexing_maps for the operands of the fused operation are same as
1166       // those for the operands of the linalgOp without the indexing map at
1167       // operand.index()
1168       SmallVector<AffineMap, 4> fusedIndexMaps = llvm::to_vector<4>(
1169           linalgOp.indexing_maps().getAsValueRange<AffineMapAttr>());
1170       fusedIndexMaps.erase(std::next(fusedIndexMaps.begin(), operand.index()));
1171 
1172       // Check if the operation shapes to loops map is computable.
1173       if (!inversePermutation(concatAffineMaps(fusedIndexMaps))) {
1174         return rewriter.notifyMatchFailure(
1175             linalgOp, "fused op loop bound computation failed");
1176       }
1177 
1178       // The operands list is same as the linalgOp with the argument for
1179       // constant index dropped.
1180       SmallVector<Value, 4> fusedOperands(linalgOp.getInputs());
1181       fusedOperands.erase(std::next(fusedOperands.begin(), operand.index()));
1182 
1183       // Create a constant scalar value from the splat constant.
1184       Value scalarConstant = rewriter.create<ConstantOp>(
1185           constantOp.getLoc(),
1186           constantOp.value().cast<DenseElementsAttr>().getSplatValue());
1187 
1188       LinalgOp fusedOp = createLinalgOpOfSameType(
1189           linalgOp, rewriter, rewriter.getUnknownLoc(),
1190           linalgOp->getResultTypes(),
1191           /*inputs=*/fusedOperands,
1192           /*outputs=*/linalgOp.getOutputs(),
1193           rewriter.getAffineMapArrayAttr(fusedIndexMaps),
1194           linalgOp.iterator_types(),
1195           /*doc=*/nullptr,
1196           /*library_call=*/nullptr,
1197           /*sparse=*/nullptr);
1198 
1199       // Map the block argument corresponding to the replaced argument with the
1200       // scalar constant.
1201       Region &linalgOpRegion = linalgOp->getRegion(0);
1202       Block &entryBlock = *linalgOpRegion.begin();
1203       unsigned argIndex = entryBlock.getNumArguments() -
1204                           linalgOp.getNumShapedOperands() + operand.index();
1205       BlockAndValueMapping mapping;
1206       mapping.map(entryBlock.getArgument(argIndex), scalarConstant);
1207       Region &fusedRegion = fusedOp->getRegion(0);
1208       rewriter.cloneRegionBefore(linalgOpRegion, fusedRegion,
1209                                  fusedRegion.begin(), mapping);
1210       rewriter.replaceOp(linalgOp, fusedOp->getResults());
1211       return success();
1212     }
1213     return failure();
1214   }
1215 
1216 private:
1217   ControlElementwiseOpsFusionFn controlFn;
1218 };
1219 } // namespace
1220 
1221 static Optional<SmallVector<Value, 1>>
1222 fuseElementwiseOps(PatternRewriter &rewriter, OpOperand &consumerOpOperand,
1223                    const ControlElementwiseOpsFusionFn &controlFn) {
1224   Operation *producer = consumerOpOperand.get().getDefiningOp();
1225   if (!producer || producer->getNumResults() != 1)
1226     return llvm::None;
1227 
1228   // Fuse when consumer is GenericOp or IndexedGenericOp.
1229   if (!isa<GenericOp, IndexedGenericOp>(consumerOpOperand.getOwner()) ||
1230       !isa<GenericOp, IndexedGenericOp>(producer))
1231     return llvm::None;
1232 
1233   return fuseElementwiseOpsImpl(cast<LinalgOp>(producer), consumerOpOperand,
1234                                 controlFn, rewriter);
1235 }
1236 
1237 namespace {
1238 /// Patterns to fuse a generic op, with the producer of its operands.
1239 template <typename LinalgOpTy>
1240 class FuseElementwiseOps : public OpRewritePattern<LinalgOpTy> {
1241 public:
1242   FuseElementwiseOps(MLIRContext *context, ControlElementwiseOpsFusionFn &fun,
1243                      PatternBenefit benefit = 1)
1244       : OpRewritePattern<LinalgOpTy>(context, benefit), controlFn(fun) {}
1245 
1246   LogicalResult matchAndRewrite(LinalgOpTy op,
1247                                 PatternRewriter &rewriter) const override {
1248     // Find the first operand that is defined by another generic op on tensors.
1249     for (OpOperand &opOperand : op.getShapedOpOperands()) {
1250       LinalgOp producerOp =
1251           dyn_cast_or_null<LinalgOp>(opOperand.get().getDefiningOp());
1252       if (!producerOp || !producerOp.hasTensorSemantics())
1253         continue;
1254       Optional<SmallVector<Value, 1>> fusedOpResults =
1255           fuseElementwiseOps(rewriter, opOperand, controlFn);
1256       if (fusedOpResults) {
1257         rewriter.replaceOp(op, *fusedOpResults);
1258         return success();
1259       }
1260     }
1261     return failure();
1262   }
1263 
1264 private:
1265   ControlElementwiseOpsFusionFn controlFn;
1266 };
1267 
1268 /// Pass that fuses generic ops on tensors. Used only for testing.
1269 struct FusionOfTensorOpsPass
1270     : public LinalgFusionOfTensorOpsBase<FusionOfTensorOpsPass> {
1271   void runOnOperation() override {
1272     Operation *op = getOperation();
1273     RewritePatternSet patterns(op->getContext());
1274     populateElementwiseOpsFusionPatterns(
1275         patterns,
1276         LinalgElementwiseFusionOptions().setAllowFoldingUnitDimReshapes(
1277             allowFoldingUnitDimReshapes));
1278     (void)applyPatternsAndFoldGreedily(op->getRegions(), std::move(patterns));
1279   }
1280 };
1281 
1282 /// Pass to test folding of reshape op with generic/indexed_generic ops by
1283 /// linearization.
1284 struct FoldReshapeOpsByLinearizationPass
1285     : public LinalgFoldReshapeOpsByLinearizationBase<
1286           FoldReshapeOpsByLinearizationPass> {
1287   void runOnOperation() override {
1288     Operation *op = getOperation();
1289     RewritePatternSet patterns(op->getContext());
1290     populateFoldReshapeOpsByLinearizationPatterns(patterns);
1291     (void)applyPatternsAndFoldGreedily(op->getRegions(), std::move(patterns));
1292   }
1293 };
1294 
1295 } // namespace
1296 
1297 void mlir::linalg::populateFoldReshapeOpsByLinearizationPatterns(
1298     RewritePatternSet &patterns) {
1299   patterns.add<FoldProducerReshapeOpByLinearization<GenericOp, false>,
1300                FoldProducerReshapeOpByLinearization<IndexedGenericOp, false>,
1301                FoldConsumerReshapeOpByLinearization<false>>(
1302       patterns.getContext());
1303 }
1304 
1305 void mlir::linalg::populateFoldUnitDimsReshapeOpsByLinearizationPatterns(
1306     RewritePatternSet &patterns) {
1307   patterns.add<FoldProducerReshapeOpByLinearization<GenericOp, true>,
1308                FoldProducerReshapeOpByLinearization<IndexedGenericOp, true>,
1309                FoldConsumerReshapeOpByLinearization<true>>(
1310       patterns.getContext());
1311 }
1312 
1313 void mlir::linalg::populateFoldReshapeOpsByExpansionPatterns(
1314     RewritePatternSet &patterns, bool allowFoldingUnitDimReshapes) {
1315   patterns.add<FoldReshapeWithGenericOpByExpansion>(patterns.getContext());
1316   patterns.add<FoldWithProducerReshapeOpByExpansion<GenericOp>,
1317                FoldWithProducerReshapeOpByExpansion<IndexedGenericOp>>(
1318       patterns.getContext(), allowFoldingUnitDimReshapes);
1319 }
1320 
1321 void mlir::linalg::populateElementwiseOpsFusionPatterns(
1322     RewritePatternSet &patterns, LinalgElementwiseFusionOptions options) {
1323   auto *context = patterns.getContext();
1324   patterns
1325       .add<FuseElementwiseOps<GenericOp>, FuseElementwiseOps<IndexedGenericOp>,
1326            FoldSplatConstants<GenericOp>, FoldSplatConstants<IndexedGenericOp>>(
1327           context, options.controlElementwiseOpsFusionFn);
1328   populateFoldReshapeOpsByExpansionPatterns(
1329       patterns, options.allowFoldingUnitDimReshapes);
1330   AffineApplyOp::getCanonicalizationPatterns(patterns, context);
1331   GenericOp::getCanonicalizationPatterns(patterns, context);
1332   IndexedGenericOp::getCanonicalizationPatterns(patterns, context);
1333   TensorReshapeOp::getCanonicalizationPatterns(patterns, context);
1334 }
1335 
1336 std::unique_ptr<Pass> mlir::createLinalgFusionOfTensorOpsPass() {
1337   return std::make_unique<FusionOfTensorOpsPass>();
1338 }
1339 
1340 std::unique_ptr<Pass> mlir::createFoldReshapeOpsByLinearizationPass() {
1341   return std::make_unique<FoldReshapeOpsByLinearizationPass>();
1342 }
1343