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 pass.
10 //
11 //===----------------------------------------------------------------------===//
12 
13 #include "PassDetail.h"
14 #include "mlir/Dialect/Affine/IR/AffineOps.h"
15 #include "mlir/Dialect/Linalg/Analysis/DependenceAnalysis.h"
16 #include "mlir/Dialect/Linalg/IR/LinalgOps.h"
17 #include "mlir/Dialect/Linalg/IR/LinalgTypes.h"
18 #include "mlir/Dialect/Linalg/Passes.h"
19 #include "mlir/Dialect/Linalg/Transforms/Transforms.h"
20 #include "mlir/Dialect/Linalg/Utils/Utils.h"
21 #include "mlir/Dialect/MemRef/IR/MemRef.h"
22 #include "mlir/Dialect/Tensor/IR/Tensor.h"
23 #include "mlir/IR/AffineExpr.h"
24 #include "mlir/IR/AffineMap.h"
25 #include "mlir/IR/Dominance.h"
26 #include "mlir/Support/LLVM.h"
27 #include "mlir/Transforms/GreedyPatternRewriteDriver.h"
28 #include "mlir/Transforms/RegionUtils.h"
29 #include "llvm/ADT/MapVector.h"
30 #include "llvm/ADT/ScopeExit.h"
31 #include "llvm/Support/CommandLine.h"
32 #include "llvm/Support/Debug.h"
33 
34 #include <set>
35 
36 #define DEBUG_TYPE "linalg-fusion"
37 
38 using namespace mlir;
39 using namespace mlir::linalg;
40 
41 using llvm::dbgs;
42 
43 /// Implements a simple high-level fusion pass on linalg structured operations.
44 ///
45 /// In each block, linalg ops are processed in reverse textual order.
46 /// Given a linalg op `O`, fusion occurs by:
47 ///   1. inspecting the linalg ops that write into the views read by `O`. There
48 ///      are 2 cases:
49 ///      a) buffer case: use the SSA value of the views and a simple alias
50 ///         analysis on subview ops to determine producer-consumer dependences;
51 ///      b) tensor case: use SSA use-def chains on subtensor ops;
52 ///   2. greedily fuse the linalg ops that produce the subview/subtensor.
53 ///   3. inspect the fused ops and determine whether they have other remaining
54 ///      LinalgOp uses. If not, then erase the original producing linalg op.
55 ///
56 /// More advanced use cases, analyses as well as profitability heuristics are
57 /// left for future work.
58 
59 struct ShapeDimension {
60   Value shape;
61   unsigned dimension;
62 };
63 
64 // Given an `op`, returns the first (`shape`, `dimension`) pair that identifies
65 // the loop range at `loopDepth`. The semantics of the loopToOperandRangesMaps
66 // guarantees at least one such dimension is found. If multiple candidates exist
67 // they must agree by construction (i.e. have the same size) and we just return
68 // the first one.
69 static ShapeDimension
70 getShapeDefiningLoopRange(LinalgOp op, unsigned loopDepth,
71                           bool fromSubViewOpOnly = false) {
72   // Iterate over the inputs and outputs in order.
73   // Extract the subranges from the linearized ranges.
74   for (OpOperand *opOperand : op.getInputAndOutputOperands()) {
75     // The method `getRangeFromOperandShape` requires using SubViewOp or
76     // SubTensorOps. If the value isnt defined from there continue.
77     // todo: The method should be adapted to get the values from
78     // `ViewInterface`. The interface needs a `getOrCreateRanges` method which
79     // currently returns a `linalg.range`. The fix here is to move this op to
80     // `std` dialect and add the method to `ViewInterface`.
81     if (fromSubViewOpOnly && !isa_and_nonnull<memref::SubViewOp, SubTensorOp>(
82                                  opOperand->get().getDefiningOp()))
83       continue;
84 
85     AffineMap map = op.getTiedIndexingMap(opOperand);
86     LLVM_DEBUG(llvm::dbgs() << "getShapeDefiningLoopRange I/O idx: "
87                             << opOperand->getOperandNumber() << "\n");
88     LLVM_DEBUG(llvm::dbgs()
89                << "getShapeDefiningLoopRange map: " << map << "\n");
90     SmallVector<Value, 8> shapeRanges(map.getNumResults(), nullptr);
91     for (auto en : llvm::enumerate(map.getResults())) {
92       auto dimExpr = en.value().dyn_cast<AffineDimExpr>();
93       if (!dimExpr)
94         continue;
95       if (loopDepth == en.value().cast<AffineDimExpr>().getPosition()) {
96         LLVM_DEBUG(llvm::dbgs() << "getShapeDefiningLoopRange loopDepth: "
97                                 << loopDepth << "\n");
98         LLVM_DEBUG(llvm::dbgs() << "getShapeDefiningLoopRange shape: "
99                                 << opOperand->get() << "\n");
100         return ShapeDimension{opOperand->get(),
101                               static_cast<unsigned>(en.index())};
102       }
103     }
104   }
105   llvm_unreachable("Expect to be able to extract a shape defining loop range");
106 }
107 
108 // Return tiled operands for the fused producer op. When fusing into
109 // `linalg.tiled_loop` one has to update `input` and `output` arguments of the
110 // loop correspondingly.
111 // Each input tensor of the producer op has to be added to `inputs` of the
112 // `tiled_loop` if it is not present there already. Each output tensor has to
113 // be added either to `inputs` or to `outputs` of `linalg.tiled_loop` depending
114 // on whether the correponding result is an input or an output to the loop.
115 //
116 // NOTE: This way of updating the arguments of the `tiled_loop` assumes that the
117 // intermediate result is not used by any other operation but the consumer. A
118 // more generic way is to append all missing output tensors of the producer to
119 // the tiled loop outputs and hence modify the number of the results, since we
120 // would need to add the intermediate results to `linalg.yield`. After that a
121 // canonicalization pass would move the unused output args of the `tiled_loop`
122 // to the `input` section.
123 static SmallVector<Value> getTiledOperands(OpBuilder &b, LinalgOp producer) {
124   auto tiledLoop = dyn_cast<TiledLoopOp>(b.getBlock()->getParentOp());
125   if (!tiledLoop)
126     return producer.getInputAndOutputOperands();
127 
128   SmallVector<Value> tiledOperands;
129   assert(producer.hasTensorSemantics() &&
130          "only fusion on tensors is currently supported for TiledLinalgOp");
131 
132   for (OpOperand *producerInput : producer.getInputTensorOperands()) {
133     OpOperand *addedInput = tiledLoop.findInputOperand(producerInput->get());
134     if (addedInput == nullptr)
135       addedInput = &tiledLoop.appendInputOperand(b, producerInput->get());
136     BlockArgument addedBlockArg = tiledLoop.getTiedBlockArgument(*addedInput);
137     tiledOperands.push_back(addedBlockArg);
138   }
139   for (OpOperand *producerOutput : producer.getOutputTensorOperands()) {
140     OpResult result = producer.getTiedOpResult(producerOutput);
141     OpOperand *resultInputOperand = tiledLoop.findInputOperand(result);
142     OpOperand *resultOutputOperand = tiledLoop.findOutputOperand(result);
143     assert((resultInputOperand != nullptr) ^ (resultOutputOperand != nullptr) &&
144            "The result should be present in `input` or `output` args of "
145            "`tiled_loop");
146 
147     bool isInput = resultInputOperand;
148     int opNumber = isInput ? resultInputOperand->getOperandNumber()
149                            : resultOutputOperand->getOperandNumber();
150 
151     OpOperand *addedOutput = tiledLoop.findOutputOperand(producerOutput->get());
152     if (addedOutput == nullptr)
153       addedOutput =
154           isInput ? &tiledLoop.appendInputOperand(b, producerOutput->get())
155                   : &tiledLoop.appendOutputOperand(b, producerOutput->get());
156 
157     OpOperand &resultOperand = tiledLoop->getOpOperand(opNumber);
158     auto addedBlockArg = tiledLoop.getTiedBlockArgument(*addedOutput);
159     auto resultOperandBlockArg = tiledLoop.getTiedBlockArgument(resultOperand);
160     resultOperandBlockArg.replaceAllUsesWith(addedBlockArg);
161     tiledLoop.eraseOperand(b, resultOperand);
162     tiledOperands.push_back(addedBlockArg);
163   }
164   return tiledOperands;
165 }
166 
167 /// Fuses the producer by cloning the `producer`. The `fusedLoopsAndRanges`
168 /// provides the loop range information for the fused loops. The rest are
169 /// obtained from the producer itself, since they are not tiled + fused.
170 static LinalgOp fuse(OpBuilder &b, LinalgOp producer,
171                      const DenseMap<unsigned, Range> &fusedLoopsAndRanges) {
172   SmallVector<Value, 8> ivs, tileSizes, sizeBounds;
173   SmallVector<Range, 8> loopRanges;
174   Location loc = producer.getLoc();
175   auto zero = b.create<ConstantIndexOp>(loc, 0);
176   auto one = b.create<ConstantIndexOp>(loc, 1);
177 
178   for (unsigned i = 0, e = producer.getNumLoops(); i < e; ++i) {
179     auto it = fusedLoopsAndRanges.find(i);
180     if (it != fusedLoopsAndRanges.end()) {
181       ivs.push_back(it->second.offset);
182       tileSizes.push_back(it->second.size);
183       sizeBounds.push_back(nullptr);
184       loopRanges.push_back(it->second);
185       LLVM_DEBUG(llvm::dbgs() << "tiled loop#" << i << " with LoopRange "
186                               << loopRanges.back() << "\n");
187     } else {
188       auto shapeDim = getShapeDefiningLoopRange(producer, i);
189       Value dim = b.createOrFold<memref::DimOp>(loc, shapeDim.shape,
190                                                 shapeDim.dimension);
191       tileSizes.push_back(zero);
192       sizeBounds.push_back(dim);
193       loopRanges.push_back(Range{zero, dim, one});
194       LLVM_DEBUG(llvm::dbgs() << "full loop#" << i << " with LoopRange "
195                               << loopRanges.back() << "\n");
196     }
197   }
198 
199   SmallVector<Value, 8> clonedShapes;
200   clonedShapes.reserve(producer.getNumInputsAndOutputs());
201 
202   // Compute subranges for all tensor input/output operands.
203   clonedShapes.append(makeTiledShapes(b, loc, producer,
204                                       getTiledOperands(b, producer), ivs,
205                                       tileSizes, sizeBounds));
206 
207   // Append the other operands.
208   auto operands = producer.getAssumedNonShapedOperands();
209   clonedShapes.append(operands.begin(), operands.end());
210 
211   // Iterate over the results in order.
212   // Extract the subtensor type from the linearized range.
213   // Since we do not enforce any canonicalizations on the fly, this is always
214   // fully dynamic at construction time.
215   SmallVector<Type, 4> resultTypes;
216   resultTypes.reserve(producer->getNumResults());
217   for (RankedTensorType t : producer.getOutputTensorTypes()) {
218     unsigned rank = t.getRank();
219     SmallVector<int64_t, 4> staticOffsetsVector(
220         rank, ShapedType::kDynamicStrideOrOffset);
221     SmallVector<int64_t, 4> staticSizesVector(rank, ShapedType::kDynamicSize);
222     SmallVector<int64_t, 4> staticStridesVector(
223         rank, ShapedType::kDynamicStrideOrOffset);
224     resultTypes.push_back(SubTensorOp::inferResultType(
225         t.cast<RankedTensorType>(), staticOffsetsVector, staticSizesVector,
226         staticStridesVector));
227   }
228 
229   Operation *clonedOp = producer.clone(b, loc, resultTypes, clonedShapes);
230   // When the producer has index semantics, we have to transform the indices of
231   // the producer according to the tiling of the consumer, i.e. offset them by
232   // the values computed in `loopRanges`.
233   assert(!isa<IndexedGenericOp>(producer) && "unexpected op");
234   if (producer.hasIndexSemantics()) {
235     assert(clonedOp->getNumRegions() == 1 &&
236            clonedOp->getRegion(0).getBlocks().size() == 1 &&
237            "expected producer to have one block.");
238     // Shift all indices by the tile offset.
239     Block &block = clonedOp->getRegion(0).front();
240     for (IndexOp indexOp : block.getOps<IndexOp>()) {
241       OpBuilder::InsertionGuard g(b);
242       b.setInsertionPointAfter(indexOp);
243       AffineExpr index, offset;
244       bindDims(b.getContext(), index, offset);
245       AffineApplyOp applyOp = b.create<AffineApplyOp>(
246           indexOp.getLoc(), index + offset,
247           ValueRange{indexOp.getResult(), loopRanges[indexOp.dim()].offset});
248       indexOp.getResult().replaceAllUsesExcept(applyOp, applyOp);
249     }
250   }
251 
252   return clonedOp;
253 }
254 
255 /// Get the loop range for a dimension `dim` based on the `shapedOperand`. It is
256 /// expected to be defined by a subview op or a subtensor op.
257 static Range getRangeFromOperandShape(OpBuilder &b, Location loc,
258                                       Value shapedOperand, unsigned dim) {
259   Operation *shapeProducingOp = shapedOperand.getDefiningOp();
260   if (auto subViewOp = dyn_cast<memref::SubViewOp>(shapeProducingOp))
261     return subViewOp.getOrCreateRanges(b, loc)[dim];
262   if (auto subTensorOp = dyn_cast<SubTensorOp>(shapeProducingOp))
263     return subTensorOp.getOrCreateRanges(b, loc)[dim];
264   llvm_unreachable("SubviewOp or SubTensorOp expected");
265 }
266 
267 /// Fuses the producer into the loop immediately enclosing the consumer.
268 /// This is achieved by "recomputing" the producer at the time it
269 /// is needed just before the consumer.
270 static LinalgOp fuse(OpBuilder &b, LinalgOp producerOp, AffineMap producerMap,
271                      OpOperand &consumerOpOperand) {
272   LLVM_DEBUG(llvm::dbgs() << "Producer map: " << producerMap << "\n");
273   DenseMap<unsigned, Range> fusedLoopsAndRanges;
274   Value shapedOperand = consumerOpOperand.get();
275   for (auto en : llvm::enumerate(producerMap.getResults())) {
276     unsigned posInProducerLoop = en.value().cast<AffineDimExpr>().getPosition();
277     fusedLoopsAndRanges[posInProducerLoop] = getRangeFromOperandShape(
278         b, consumerOpOperand.getOwner()->getLoc(), shapedOperand, en.index());
279   }
280   return fuse(b, producerOp, fusedLoopsAndRanges);
281 }
282 
283 // Encode structural fusion safety preconditions.
284 // Some of these will be lifted in the future with better analysis.
285 static bool isStructurallyFusableProducer(LinalgOp producer, Value consumedView,
286                                           LinalgOp consumer) {
287   assert(producer.hasBufferSemantics() &&
288          "expected linalg op with buffer semantics");
289   assert(consumer.hasBufferSemantics() &&
290          "expected linalg op with buffer semantics");
291   if (producer.getNumOutputs() != 1) {
292     LLVM_DEBUG(llvm::dbgs() << "\nNot structurally fusable (multi-output)");
293     return false;
294   }
295   // Only fuse when the producer block dominates.
296   DominanceInfo dom(producer.getOperation());
297   if (!dom.dominates(producer->getBlock(), consumer->getBlock())) {
298     LLVM_DEBUG(
299         llvm::dbgs()
300         << "\nNot structurally fusable (producer block does not dominate)");
301     return false;
302   }
303   return true;
304 }
305 
306 bool mlir::linalg::isProducerLastWriteOfView(const LinalgDependenceGraph &graph,
307                                              LinalgOp consumer,
308                                              Value consumedView,
309                                              LinalgOp producer) {
310   assert(producer.hasBufferSemantics() &&
311          "expected linalg op with buffer semantics");
312   assert(consumer.hasBufferSemantics() &&
313          "expected linalg op with buffer semantics");
314   // Make some simple structural checks that alleviate the need for more
315   // complex analyses.
316   if (!isStructurallyFusableProducer(producer, consumedView, consumer)) {
317     LLVM_DEBUG(llvm::dbgs() << "\n***Not static last write due to structure:\t"
318                             << *producer.getOperation());
319     return false;
320   }
321   // Check for any interleaved write to consumedView.
322   if (!graph.findCoveringWrites(producer, consumer, consumedView).empty()) {
323     LLVM_DEBUG(llvm::dbgs() << "\n***Not fusable due to interleaved write:\t"
324                             << *producer.getOperation());
325     return false;
326   }
327   return true;
328 }
329 
330 bool mlir::linalg::isFusableInto(const LinalgDependenceGraph &graph,
331                                  LinalgOp consumer, Value consumedView,
332                                  LinalgOp producer) {
333   assert(producer.hasBufferSemantics() &&
334          "expected linalg op with buffer semantics");
335   assert(consumer.hasBufferSemantics() &&
336          "expected linalg op with buffer semantics");
337   if (!isProducerLastWriteOfView(graph, consumer, consumedView, producer))
338     return false;
339   // Check for any fusion-preventing dependence to any shape read/written that
340   // would violate dependences.
341   if (!graph.findCoveringDependences(producer, consumer).empty()) {
342     LLVM_DEBUG(llvm::dbgs()
343                << "\n***Not fusable due to an interleaved dependence:\t"
344                << *producer.getOperation());
345     return false;
346   }
347   if (auto convOp = dyn_cast<linalg::ConvOp>(producer.getOperation())) {
348     // TODO: add a level of indirection to linalg.generic.
349     if (convOp.padding())
350       return false;
351   }
352   if (auto convOp = dyn_cast<linalg::ConvOp>(consumer.getOperation())) {
353     // TODO: add a level of indirection to linalg.generic.
354     if (convOp.padding())
355       return false;
356   }
357   return true;
358 }
359 
360 /// For `consumer` with buffer semantics, find the Linalg operation on buffers
361 /// that is the last writer of `consumerOpOperand`. For now the fusable
362 /// dependence is returned as an instance of the `dependenceGraph`.
363 static Optional<LinalgDependenceGraph::LinalgDependenceGraphElem>
364 findFusableProducer(OpOperand &consumerOpOperand,
365                     const LinalgDependenceGraph &dependenceGraph) {
366   LLVM_DEBUG(llvm::dbgs() << "findFusableProducer for: "
367                           << consumerOpOperand.get() << " @"
368                           << consumerOpOperand.getOperandNumber() << " in "
369                           << *consumerOpOperand.getOwner() << "\n");
370   LinalgOp consumerOp = dyn_cast<LinalgOp>(consumerOpOperand.getOwner());
371   if (!consumerOp)
372     return {};
373 
374   // Only consider RAW and WAW atm.
375   for (auto depType : {
376            LinalgDependenceGraph::DependenceType::RAW,
377            LinalgDependenceGraph::DependenceType::WAW,
378        }) {
379     LLVM_DEBUG(llvm::dbgs()
380                << "Dependencies into: " << *consumerOp.getOperation() << "\n");
381     for (auto dependence : llvm::make_filter_range(
382              dependenceGraph.getDependencesInto(consumerOp, depType),
383              [&](LinalgDependenceGraph::LinalgDependenceGraphElem elem) {
384                LLVM_DEBUG(llvm::dbgs() << "Inspect dependence btw: "
385                                        << elem.getIndexingValue() << " and "
386                                        << elem.getDependentValue() << "\n");
387                Value v = elem.getIndexingValue();
388                Optional<unsigned> operandNum =
389                    elem.getIndexingOpViewOperandNum();
390                return isa<LinalgOp>(elem.getDependentOp()) &&
391                       v == consumerOpOperand.get() && operandNum &&
392                       operandNum.getValue() ==
393                           consumerOpOperand.getOperandNumber();
394              })) {
395       // Consumer consumes this view, `isStructurallyFusableProducer` also
396       // checks whether it is a strict subview of the producer view.
397       auto producer = cast<LinalgOp>(dependence.getDependentOp());
398       LLVM_DEBUG(llvm::dbgs()
399                  << "\n"
400                  << LinalgDependenceGraph::getDependenceTypeStr(depType)
401                  << "producer: " << *dependence.getDependentOp()
402                  << " view: " << dependence.getDependentValue() << "\n");
403 
404       // If the producer and consumer have tensor semantics, the only dependence
405       // between them is through a RAW dependence and they are fusable by
406       // construction. For buffer semantics need additional checks.
407       if (producer.hasBufferSemantics() && consumerOp.hasBufferSemantics() &&
408           isFusableInto(dependenceGraph, consumerOp, consumerOpOperand.get(),
409                         producer))
410         return dependence;
411       if (producer.hasTensorSemantics() && consumerOp.hasTensorSemantics()) {
412         assert(dependence.dependenceType ==
413                LinalgDependenceGraph::DependenceType::RAW);
414         return dependence;
415       }
416     }
417   }
418   return {};
419 }
420 
421 Optional<FusionInfo>
422 mlir::linalg::fuseProducerOfBuffer(OpBuilder &b, OpOperand &consumerOpOperand,
423                                    const LinalgDependenceGraph &graph) {
424   Optional<LinalgDependenceGraph::LinalgDependenceGraphElem> fusableDependence =
425       findFusableProducer(consumerOpOperand, graph);
426   if (!fusableDependence)
427     return llvm::None;
428 
429   // Canonicalize indexed generic ops before fusion.
430   if (isa<IndexedGenericOp>(fusableDependence->getDependentOp()))
431     return llvm::None;
432 
433   LinalgOp producerOp = dyn_cast<LinalgOp>(fusableDependence->getDependentOp());
434   if (!producerOp)
435     return llvm::None;
436 
437   // If producer is already in the same block as consumer, we are done.
438   if (consumerOpOperand.get().getParentBlock() ==
439       fusableDependence->getDependentValue().getParentBlock())
440     return llvm::None;
441 
442   Optional<AffineMap> producerMap =
443       fusableDependence->getDependentOpViewIndexingMap();
444   if (!producerMap)
445     return llvm::None;
446 
447   // Must be a subview or a slice to guarantee there are loops we can fuse
448   // into.
449   auto subView = consumerOpOperand.get().getDefiningOp<memref::SubViewOp>();
450   if (!subView) {
451     LLVM_DEBUG(llvm::dbgs() << "\nNot fusable (not a subview)");
452     return llvm::None;
453   }
454 
455   // Fuse `producer` just before `consumer`.
456   OpBuilder::InsertionGuard g(b);
457   b.setInsertionPoint(consumerOpOperand.getOwner());
458   LLVM_DEBUG(llvm::dbgs() << "Fuse into consumer: "
459                           << *consumerOpOperand.getOwner() << "\n");
460 
461   auto fusedProducer = fuse(b, producerOp, *producerMap, consumerOpOperand);
462   return FusionInfo{producerOp, fusedProducer};
463 }
464 
465 /// Walk back use-def chain through scf::For yields.
466 /// Sets `producer` and `outputIndex` if it finds a producer LinalgOp
467 
468 // TODO(ravishankarm, ntv): This can be moved into the dependence graphs
469 // dependence tracking since the dependence tracking is similar to what is done
470 // w.r.t to buffers.
471 static void getProducerOfTensor(Value tensor, OpResult &opResult) {
472   if (!tensor.getType().isa<RankedTensorType>())
473     return;
474 
475   while (true) {
476     LLVM_DEBUG(llvm::dbgs() << "\ngetProducerOfTensor: " << tensor);
477     if (auto linalgOp = tensor.getDefiningOp<LinalgOp>()) {
478       opResult = tensor.cast<OpResult>();
479       return;
480     }
481     if (auto subTensorOp = tensor.getDefiningOp<SubTensorOp>()) {
482       tensor = subTensorOp.source();
483       continue;
484     }
485     if (auto blockArg = tensor.dyn_cast<BlockArgument>()) {
486       if (auto forOp = blockArg.getDefiningOp<scf::ForOp>()) {
487         tensor = *(forOp.getIterOperands().begin() + blockArg.getArgNumber());
488         continue;
489       }
490     }
491     return;
492   }
493 }
494 
495 Optional<FusionInfo>
496 mlir::linalg::fuseProducerOfTensor(OpBuilder &b, OpOperand &consumerOpOperand) {
497   Value inputTensor = consumerOpOperand.get();
498   OpResult producerOpResult;
499   getProducerOfTensor(inputTensor, producerOpResult);
500   if (!producerOpResult) {
501     LLVM_DEBUG(llvm::dbgs() << "\nUnable to find producer");
502     return {};
503   }
504   return fuseProducerOfTensor(b, producerOpResult, consumerOpOperand);
505 }
506 
507 Optional<FusionInfo>
508 mlir::linalg::fuseProducerOfTensor(OpBuilder &b, OpResult producerOpResult,
509                                    OpOperand &consumerOpOperand) {
510   // Canonicalize indexed generic ops before fusion.
511   if (isa<IndexedGenericOp>(producerOpResult.getOwner()))
512     return llvm::None;
513 
514   auto producerOp = dyn_cast<LinalgOp>(producerOpResult.getOwner());
515   if (!producerOp)
516     return llvm::None;
517 
518   LinalgOp consumerOp = dyn_cast<LinalgOp>(consumerOpOperand.getOwner());
519   if (!consumerOp)
520     return llvm::None;
521 
522   Value inputTensor = consumerOpOperand.get();
523 
524   // Must be a subtensor to guarantee there are loops we can fuse into.
525   auto subTensor = inputTensor.getDefiningOp<SubTensorOp>();
526   if (!subTensor) {
527     LLVM_DEBUG(llvm::dbgs()
528                << "\nNot fusable, not a subtensor: " << inputTensor);
529     return {};
530   }
531 
532   // If producer is already in the same block as consumer, we are done.
533   if (consumerOpOperand.get().getParentBlock() ==
534       producerOpResult.getParentBlock())
535     return {};
536 
537   // Insert fused `producer` just before `consumer`.
538   OpBuilder::InsertionGuard g(b);
539   b.setInsertionPoint(consumerOp);
540   LLVM_DEBUG(llvm::dbgs() << "Fuse into consumer: " << *consumerOp << "\n");
541   OpOperand *opOperand =
542       producerOp.getOutputOperand(producerOpResult.getResultNumber());
543   LinalgOp fusedProducer =
544       fuse(b, producerOp, producerOp.getTiedIndexingMap(opOperand),
545            consumerOpOperand);
546 
547   // Replace use.
548   // Canonicalizations are not guaranteed to have happened before constructing
549   // `fusedProducer`. In the tensor case this can result in temporary type
550   // mismatches. Insert a `tensor.cast` op to propagate the transformation
551   // invariant that types are compatible.
552   Value def = fusedProducer->getResult(producerOpResult.getResultNumber());
553   Type consumerType = consumerOpOperand.get().getType();
554   if (consumerType != def.getType())
555     def = b.create<tensor::CastOp>(fusedProducer.getLoc(), consumerType, def);
556   consumerOpOperand.set(def);
557   return FusionInfo{cast<LinalgOp>(producerOpResult.getOwner()), fusedProducer};
558 }
559 
560 /// Prune all dimensions that are of reduction iterator type from `map`.
561 static AffineMap pruneReductionDimsFromMap(ArrayRef<Attribute> iteratorTypes,
562                                            AffineMap map) {
563   llvm::SmallDenseSet<unsigned> projectedDims;
564   for (auto attr : llvm::enumerate(iteratorTypes)) {
565     if (!isParallelIterator(attr.value()))
566       projectedDims.insert(attr.index());
567   }
568   return getProjectedMap(map, projectedDims);
569 }
570 
571 /// Returns the mapping from iterations in the consumer that write to the same
572 /// location as the iterations in the producer. To do so use
573 /// - indexing map of the fused view in the consumer : consumerIndexMap
574 /// - indexing map of the fused view in the producer : producerIndexMap
575 ///     consumerLoopToProducerLoop =
576 ///       inverse(producerIndexMap).compose(consumerIndexMap)
577 static Optional<AffineMap> getConsumerLoopToProducerLoopMap(
578     LinalgDependenceGraph::LinalgDependenceGraphElem dependence) {
579   auto producer = dyn_cast<LinalgOp>(dependence.getDependentOp());
580   if (!producer)
581     return None;
582 
583   Optional<AffineMap> producerIndexingMap =
584       dependence.getDependentOpViewIndexingMap();
585   Optional<AffineMap> consumerIndexingMap =
586       dependence.getIndexingOpViewIndexingMap();
587   if (!producerIndexingMap || !consumerIndexingMap)
588     return None;
589 
590   AffineMap prunedProducerIndexingMap = pruneReductionDimsFromMap(
591       producer.iterator_types().getValue(), *producerIndexingMap);
592   if (!prunedProducerIndexingMap.isPermutation())
593     return None;
594 
595   if (consumerIndexingMap->getNumResults() !=
596       prunedProducerIndexingMap.getNumResults())
597     return None;
598 
599   LLVM_DEBUG({
600     llvm::dbgs() << "\t producerMap : ";
601     producerIndexingMap->print(llvm::dbgs());
602     llvm::dbgs() << "  pruned : ";
603     prunedProducerIndexingMap.print(llvm::dbgs());
604     llvm::dbgs() << "\n";
605     llvm::dbgs() << "\t consumerMap : ";
606     consumerIndexingMap->print(llvm::dbgs());
607     llvm::dbgs() << "\n";
608   });
609 
610   AffineMap invProducerIndexMap = inversePermutation(prunedProducerIndexingMap);
611   if (!invProducerIndexMap)
612     return None;
613 
614   return invProducerIndexMap.compose(*consumerIndexingMap);
615 }
616 
617 /// Given a projected permutation `map`, returns true if the map changes the
618 /// order in which the fused loop dimension appear.
619 static bool doesTransposeAccess(AffineMap map,
620                                 const std::set<unsigned> &fusableLoops) {
621   Optional<unsigned> lastFusableLoop;
622   for (unsigned pos : llvm::map_range(map.getResults(), [](AffineExpr expr) {
623          return expr.cast<AffineDimExpr>().getPosition();
624        })) {
625     if (!fusableLoops.count(pos))
626       continue;
627     if (!lastFusableLoop) {
628       lastFusableLoop = pos;
629       continue;
630     }
631     if (pos <= lastFusableLoop.getValue())
632       return true;
633     lastFusableLoop = pos;
634   }
635   return false;
636 }
637 
638 /// Returns the positions of the loop in `op` that can be tiled based on the
639 /// operations that are to be fused with it. For example, in a
640 ///
641 ///   linalg.matmul ins(%a, %b : ...) outs(%c : ...)
642 ///
643 /// if the producer of %a needs to be fused with this op, only the `i` loop of
644 /// the matmul can be tiled while fusing. If producer of %a, and %b are to be
645 /// fused, then no loops can be tiled while fusing. The conditions used are:
646 /// 1. Only parallel loops can be used for tile + fuse. Find the number of
647 ///    common outer parallel loops between the op and its producers being fused.
648 /// 2. Of the parallel loops only some can be fused. Only those loops can be
649 ///    fused such where the fusable loops iteration space only touches one tile
650 ///    of the fused operation. This is because the producer (which is writing
651 ///    the fused subview) has update semantics.
652 ///
653 /// Since an inverse computation is needed, we need to consider the projection
654 /// of the producerIndexMap w.r.t the parallel loops.  The actual fusable loops
655 /// are the dimensions of the consumerLoopToProducerLoop map that correspond to
656 /// parallel loops and appear in the result of the map
657 ///
658 /// Example 1:
659 ///   linalg.fill(%c, %cst)
660 ///   linalg.matmul ins(%a, %b) outs(%c)
661 ///     Number of parallel loops : 2
662 ///     producerIndexMap = affine_map<(i, j) ->(i , j)>
663 ///     consumerIndexMap = affine_map<(i, j, k) -> (i, j)>
664 ///     consumerLoopToProducerLoop = affine_map<(i, j, k) -> (i, j)>
665 ///     Fused dimensions : i, j
666 ///
667 /// Example 2:
668 ///   linalg.matmul ins(%a, %b) outs(%c)
669 ///   linalg.generic {indexing_maps = [affine_map<(i, j) -> (j, i)>, ...
670 ///                   iterator_types = ["parallel", "parallel"]}
671 ///     ins(%c) ...
672 ///
673 ///     Number of parallel loops = 2:
674 ///     producerIndexMap (projected to parallel loops) =
675 ///       affine_map<(i, j) -> (i, j)>
676 ///     consumerLoopToProducerLoop2 = affine_map<(i, j) -> (j, i)>
677 ///     Fused dimensions : i, j
678 ///
679 /// Example 3:
680 ///   linalg.copy(%s, %b)
681 ///   linalg.matmul ins(%a, %b) outs(%c)
682 ///
683 ///   Number of parallel loops = 2
684 ///   produceIndexMap : affine_map<(i, j) -> (i, j)>
685 ///   consumerLoopToProduceLoops = affine_map<(i, j, k) -> (k, j)>
686 ///     submap with only parallel loops = affine_map<(i, j) -> (j)>
687 ///   Fused dimensions : j
688 static std::set<unsigned>
689 collectFusableLoops(ArrayRef<LinalgOp> ops,
690                     const FusableOpDependencesTy &fusableDependences) {
691   assert(!ops.empty());
692   auto getNumOuterParallelLoops = [](LinalgOp linalgOp) {
693     return linalgOp.iterator_types()
694         .getValue()
695         .take_while([](Attribute attr) -> bool {
696           return attr.cast<StringAttr>().getValue() ==
697                  getParallelIteratorTypeName();
698         })
699         .size();
700   };
701 
702   size_t numOuterParallelLoops = getNumOuterParallelLoops(ops.back());
703   for (auto op : ops.drop_back()) {
704     numOuterParallelLoops =
705         std::min(numOuterParallelLoops, getNumOuterParallelLoops(op));
706   }
707 
708   std::set<unsigned> fusableLoops;
709   auto range = llvm::seq<unsigned>(0, numOuterParallelLoops);
710   fusableLoops.insert(range.begin(), range.end());
711 
712   for (auto op : reverse(ops)) {
713     for (auto dependence : fusableDependences.lookup(op)) {
714       LLVM_DEBUG({
715         llvm::dbgs() << "\t fusable :";
716         for (unsigned i : fusableLoops)
717           llvm::dbgs() << " " << i;
718         llvm::dbgs() << "\n";
719       });
720 
721       Optional<AffineMap> consumerLoopToProducerLoop =
722           getConsumerLoopToProducerLoopMap(dependence);
723       if (!consumerLoopToProducerLoop) {
724         op.emitRemark("failed to get map from consumer loop to producer loop");
725         return {};
726       }
727       // todo: This condition is only an implementation limitation. When fusing
728       // the operation, if the accesses in the producer/consumer are transposes
729       // of each other, the loop bounds for the tiled producer can be
730       // manipulated accordingly. This requires some additional bookkeeping in
731       // the implementation of tile+fuse that is deferred to later.
732       if (doesTransposeAccess(*consumerLoopToProducerLoop, fusableLoops)) {
733         op.emitRemark("unhandled fusion when fusion requires permutation");
734         return {};
735       }
736 
737       std::set<unsigned> candidates;
738       for (AffineExpr expr : consumerLoopToProducerLoop->getResults()) {
739         unsigned position = expr.cast<AffineDimExpr>().getPosition();
740         if (fusableLoops.count(position))
741           candidates.insert(position);
742       }
743       LLVM_DEBUG({
744         llvm::dbgs() << "\t candidates :";
745         for (unsigned i : candidates)
746           llvm::dbgs() << " " << i;
747         llvm::dbgs() << "\n";
748       });
749       if (candidates.empty())
750         return {};
751       std::swap(candidates, fusableLoops);
752     }
753   }
754 
755   return fusableLoops;
756 }
757 
758 /// Find all dependences that are fusable.
759 FusableOpDependencesTy mlir::linalg::findAllFusableDependences(
760     ArrayRef<LinalgOp> ops, const LinalgDependenceGraph &dependenceGraph) {
761   FusableOpDependencesTy fusableDependences;
762   DenseMap<Operation *, SmallVector<AffineMap, 1>> fusedProducerIndexingMap;
763   for (LinalgOp op : reverse(ops)) {
764     for (OpOperand *opOperand : op.getInputAndOutputOperands()) {
765       Optional<LinalgDependenceGraph::LinalgDependenceGraphElem>
766           fusableDependence = findFusableProducer(*opOperand, dependenceGraph);
767       if (!fusableDependence)
768         continue;
769       // Canonicalize indexed generic ops before fusion.
770       if (isa<IndexedGenericOp>(fusableDependence->getDependentOp()))
771         continue;
772       LinalgOp producerOp =
773           dyn_cast<LinalgOp>(fusableDependence->getDependentOp());
774       if (!producerOp)
775         continue;
776       // Do not fuse dependences that are to operations not in the same basic
777       // block. This avoid moving fused operations across loops that might
778       // themselves carry dependency making the fusion illegal.
779       if (producerOp->getBlock() != op->getBlock())
780         continue;
781 
782       // Make sure that the indexing map of the view used for fusion in the
783       // producer is a projected permutation.
784       Optional<AffineMap> producerMap =
785           fusableDependence->getDependentOpViewIndexingMap();
786       Optional<AffineMap> consumerMap =
787           fusableDependence->getIndexingOpViewIndexingMap();
788       assert(
789           consumerMap &&
790           "unable to find indexing map of operand/result of indexing OpView");
791       fusedProducerIndexingMap[producerOp.getOperation()].push_back(
792           *consumerMap);
793       if (!producerMap || !producerMap->isProjectedPermutation() ||
794           !consumerMap->isProjectedPermutation())
795         continue;
796 
797       fusableDependences[producerOp.getOperation()].push_back(
798           *fusableDependence);
799     }
800   }
801   // TODO: Currently fusion would not be legal if the fusable dependence is to
802   // the same producer but different indexing map in the consumer. Fix this, but
803   // in the meanwhile disallow such a fusion.
804   for (auto useIndexingMapsList : fusedProducerIndexingMap) {
805     AffineMap map1 = useIndexingMapsList.second.front();
806     for (AffineMap map2 :
807          ArrayRef<AffineMap>(useIndexingMapsList.second).drop_front()) {
808       if (map1 != map2) {
809         fusableDependences.erase(useIndexingMapsList.first);
810         break;
811       }
812     }
813   }
814   return fusableDependences;
815 }
816 
817 /// Tile the fused loops in the root operation, by setting the tile sizes for
818 /// all other loops to zero (those will be tiled later).
819 static Optional<TiledLinalgOp>
820 tileRootOperation(OpBuilder &b, LinalgOp op, ArrayRef<Value> tileSizeVector,
821                   const LinalgTilingOptions &options,
822                   const std::set<unsigned> &fusedLoops) {
823   SmallVector<Value, 4> tileSizes(tileSizeVector.begin(), tileSizeVector.end());
824   auto zero = b.create<ConstantIndexOp>(op.getLoc(), 0);
825   for (unsigned i = 0, e = tileSizes.size(); i != e; ++i)
826     if (!fusedLoops.count(i))
827       tileSizes[i] = zero;
828   LinalgTilingOptions tileFusedLoopsOptions = options;
829   tileFusedLoopsOptions.setTileSizes(tileSizes);
830   return tileLinalgOp(b, op, tileFusedLoopsOptions);
831 }
832 
833 /// Fuse the operations in `fusionCandidates` with `tiledOp`. Latter is expected
834 /// to be a tiled operation such that it is valid to fuse all operations in
835 /// `fusionCandidates`, i.e. move the operation within the inter-tile loops of
836 /// `tiledOp`.
837 static SmallVector<LinalgOp, 1>
838 fuseOperations(OpBuilder &b, LinalgOp rootOp, TiledLinalgOp tiledLinalgOp,
839                ArrayRef<LinalgOp> fusionCandidates,
840                const FusableOpDependencesTy &fusableDependences,
841                const std::set<unsigned> &fusedLoops) {
842   LinalgOp tiledOp = tiledLinalgOp.op;
843   OpBuilder::InsertionGuard guard(b);
844   b.setInsertionPoint(tiledOp);
845 
846   DenseMap<unsigned, Range> fusedLoopsAndRanges;
847   for (unsigned loop : fusedLoops) {
848     ShapeDimension shapeDim = getShapeDefiningLoopRange(tiledOp, loop, true);
849     fusedLoopsAndRanges[loop] = getRangeFromOperandShape(
850         b, tiledOp.getLoc(), shapeDim.shape, shapeDim.dimension);
851   }
852 
853   SmallVector<LinalgOp, 1> fusedOps(fusionCandidates.size());
854   DenseMap<Operation *, LinalgOp> origOpToFusedOp;
855   origOpToFusedOp[rootOp.getOperation()] = tiledOp;
856   for (auto candidate : enumerate(llvm::reverse(fusionCandidates))) {
857     LinalgOp origOp = candidate.value();
858     LinalgOp fusedOp = fuse(b, origOp, fusedLoopsAndRanges);
859     origOpToFusedOp[origOp.getOperation()] = fusedOp;
860     fusedOps[fusionCandidates.size() - candidate.index() - 1] = fusedOp;
861 
862     // Prepare the builder for the next insertion point.
863     auto guard = llvm::make_scope_exit([&]() { b.setInsertionPoint(fusedOp); });
864     if (!origOp.hasTensorSemantics())
865       continue;
866 
867     // If the producer consumer operations are linalg operations on tensors, the
868     // dependence is due to value produced (as a return tensor) by the producer
869     // and used in the consumer. The returned value of the fused op needs to be
870     // made the operand of the tiled/fused consumer operation. By construction
871     // the value returned by the producer is the value used by the consumer.
872     for (auto &dependence : fusableDependences.lookup(origOp.getOperation())) {
873       if (dependence.dependenceType !=
874           LinalgDependenceGraph::DependenceType::RAW)
875         continue;
876 
877       unsigned resultIndex =
878           dependence.getDependentOpViewResultNum().getValue();
879       LinalgOp consumer = origOpToFusedOp.lookup(dependence.getIndexingOp());
880       if (!consumer)
881         continue;
882 
883       Value replacementValue = fusedOp.getOperation()->getResult(resultIndex);
884       consumer.getOperation()->setOperand(
885           dependence.getIndexingOpViewOperandNum().getValue(),
886           replacementValue);
887     }
888 
889     // At this point, all Linalg uses of the tensors produced by `origOp` have
890     // been replaced. However, there may still be "output tensor"-like uses
891     // coming from WAW dependencies.
892     // All these uses are iter_args of the outermost loop (TODO: add a check).
893     // Such iter_args uses serve 2 purposes:
894     //  1. give a shape to the output
895     //  2. encode destructive updates that may be inplaceable by bufferization.
896     // To keep the second type of information while letting the unfused op die
897     // unused, we need to forward the producer output operand.
898     if (auto forOp = dyn_cast<scf::ForOp>(tiledLinalgOp.loops.front())) {
899       for (auto &operand : forOp.getIterOpOperands()) {
900         if (auto opResult = operand.get().dyn_cast<OpResult>()) {
901           if (opResult.getOwner() == origOp) {
902             Value output =
903                 origOp.getOutputOperand(opResult.getResultNumber())->get();
904             assert(output.getType().isa<RankedTensorType>());
905             operand.set(output);
906           }
907         }
908       }
909     }
910   }
911   return fusedOps;
912 }
913 
914 static Optional<TiledAndFusedLinalgOps>
915 tileAndFuseLinalgOpsImpl(OpBuilder &b, ArrayRef<LinalgOp> ops,
916                          const LinalgDependenceGraph &dependenceGraph,
917                          const LinalgTilingOptions &tilingOptions) {
918   if (ops.size() < 2)
919     return llvm::None;
920   LinalgOp rootOp = ops.back();
921   if (!llvm::all_of(
922           ops,
923           [](LinalgOp linalgOp) { return linalgOp.hasBufferSemantics(); }) &&
924       !llvm::all_of(ops, [](LinalgOp linalgOp) {
925         return linalgOp.hasTensorSemantics();
926       })) {
927     rootOp.emitError(
928         "unable to fuse operations that have tensor semantics with operations "
929         "that have buffer semantics and viceversa.");
930     return llvm::None;
931   }
932   // TODO: Support interchange with tile + fuse. This might actually help do
933   // better fusion.
934   if (!tilingOptions.interchangeVector.empty()) {
935     rootOp.emitRemark("unable to handle tile and fuse with interchange");
936     return llvm::None;
937   }
938 
939   OpBuilder::InsertionGuard guard(b);
940   b.setInsertionPoint(rootOp);
941 
942   // Find all the producers.
943   LLVM_DEBUG(llvm::dbgs() << "findAllFusableDependences\n");
944   FusableOpDependencesTy fusableDependences =
945       findAllFusableDependences(ops, dependenceGraph);
946   if (fusableDependences.empty()) {
947     LLVM_DEBUG(llvm::dbgs() << "no fusable dependencies found\n");
948     return llvm::None;
949   }
950 
951   TiledAndFusedLinalgOps ret;
952   // Find the loops that can be tiled and fused.
953   LLVM_DEBUG(llvm::dbgs() << "collectFusableLoops\n");
954   ret.fusedLoopDims = collectFusableLoops(ops, fusableDependences);
955 
956   // If there are no fusable dependences or there are no tile+fusable loops,
957   // just return.
958   if (ret.fusedLoopDims.empty()) {
959     LLVM_DEBUG(llvm::dbgs() << "no fusable loops found\n");
960     return llvm::None;
961   }
962 
963   // Tile the fused loops in the last operation in the list.
964   SmallVector<Value, 4> tileSizeVector =
965       tilingOptions.tileSizeComputationFunction(b, rootOp);
966   Optional<TiledLinalgOp> tiledRootOp = tileRootOperation(
967       b, rootOp, tileSizeVector, tilingOptions, ret.fusedLoopDims);
968   if (!tiledRootOp) {
969     rootOp.emitRemark("failed to tile the fused loops");
970     return llvm::None;
971   }
972   ret.op = tiledRootOp->op;
973   ret.fusedLoops.assign(tiledRootOp->loops.begin(), tiledRootOp->loops.end());
974 
975   // Fuse the other operations into the fused inter-tile loops produced above.
976   ret.fusedProducers = fuseOperations(b, rootOp, *tiledRootOp, ops.drop_back(),
977                                       fusableDependences, ret.fusedLoopDims);
978 
979   return ret;
980 }
981 
982 Optional<TiledAndFusedLinalgOps>
983 mlir::linalg::tileAndFuseLinalgOps(OpBuilder &b, ArrayRef<LinalgOp> ops,
984                                    const LinalgDependenceGraph &dependenceGraph,
985                                    const LinalgTilingOptions &tilingOptions) {
986   switch (tilingOptions.loopType) {
987   case LinalgTilingLoopType::Loops:
988   case LinalgTilingLoopType::ParallelLoops:
989   case LinalgTilingLoopType::TiledLoops:
990     return tileAndFuseLinalgOpsImpl(b, ops, dependenceGraph, tilingOptions);
991   default:;
992   }
993   return llvm::None;
994 }
995