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