1 //===- DependenceAnalysis.cpp - Dependence analysis on SSA views ----------===//
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 view-based alias and dependence analyses.
10 //
11 //===----------------------------------------------------------------------===//
12 
13 #include "mlir/Dialect/Linalg/Analysis/DependenceAnalysis.h"
14 #include "mlir/Dialect/Linalg/IR/LinalgOps.h"
15 #include "mlir/Dialect/StandardOps/IR/Ops.h"
16 
17 #include "llvm/Support/CommandLine.h"
18 #include "llvm/Support/Debug.h"
19 
20 #define DEBUG_TYPE "linalg-dependence-analysis"
21 
22 using namespace mlir;
23 using namespace mlir::linalg;
24 
25 using llvm::dbgs;
26 
27 Value Aliases::find(Value v) {
28   if (v.isa<BlockArgument>())
29     return v;
30 
31   auto it = aliases.find(v);
32   if (it != aliases.end()) {
33     assert(it->getSecond().getType().isa<BaseMemRefType>() &&
34            "Memref expected");
35     return it->getSecond();
36   }
37 
38   while (true) {
39     if (v.isa<BlockArgument>())
40       return v;
41 
42     Operation *defOp = v.getDefiningOp();
43     if (!defOp)
44       return v;
45 
46     // Treat RegionBranchOpInterfaces like an allocate and don't try to follow
47     // the aliasing further.
48     if (isa<RegionBranchOpInterface>(defOp))
49       return v;
50     if (isa<TensorToMemrefOp>(defOp))
51       return v;
52 
53     if (auto memEffect = dyn_cast<MemoryEffectOpInterface>(defOp)) {
54       // Collect all memory effects on `v`.
55       SmallVector<MemoryEffects::EffectInstance, 1> effects;
56       memEffect.getEffectsOnValue(v, effects);
57 
58       // If we have the 'Allocate' memory effect on `v`, then `v` should be the
59       // original buffer.
60       if (llvm::any_of(
61               effects, [](const MemoryEffects::EffectInstance &instance) {
62                 return isa<MemoryEffects::Allocate>(instance.getEffect());
63               }))
64         return v;
65     }
66 
67     if (auto viewLikeOp = dyn_cast<ViewLikeOpInterface>(defOp)) {
68       auto it =
69           aliases.insert(std::make_pair(v, find(viewLikeOp.getViewSource())));
70       return it.first->second;
71     }
72 
73     llvm::errs() << "View alias analysis reduces to: " << v << "\n";
74     llvm_unreachable("unsupported view alias case");
75   }
76 }
77 
78 StringRef LinalgDependenceGraph::getDependenceTypeStr(DependenceType depType) {
79   switch (depType) {
80   case LinalgDependenceGraph::DependenceType::RAW:
81     return "RAW";
82   case LinalgDependenceGraph::DependenceType::RAR:
83     return "RAR";
84   case LinalgDependenceGraph::DependenceType::WAR:
85     return "WAR";
86   case LinalgDependenceGraph::DependenceType::WAW:
87     return "WAW";
88   default:
89     break;
90   }
91   llvm_unreachable("Unexpected DependenceType");
92 }
93 
94 LinalgDependenceGraph
95 LinalgDependenceGraph::buildDependenceGraph(Aliases &aliases, FuncOp f) {
96   SmallVector<LinalgOp, 8> linalgOps;
97   f.walk([&](LinalgOp op) { linalgOps.push_back(op); });
98   return LinalgDependenceGraph(aliases, linalgOps);
99 }
100 
101 LinalgDependenceGraph::LinalgDependenceGraph(Aliases &aliases,
102                                              ArrayRef<LinalgOp> ops)
103     : aliases(aliases), linalgOps(ops.begin(), ops.end()) {
104   for (auto en : llvm::enumerate(linalgOps)) {
105     linalgOpPositions.insert(
106         std::make_pair(en.value().getOperation(), en.index()));
107   }
108   for (unsigned i = 0, e = ops.size(); i < e; ++i) {
109     for (unsigned j = i + 1; j < e; ++j) {
110       addDependencesBetween(ops[i], ops[j]);
111     }
112   }
113 }
114 
115 void LinalgDependenceGraph::addDependenceElem(DependenceType dt,
116                                               LinalgOpView indexingOpView,
117                                               LinalgOpView dependentOpView) {
118   LLVM_DEBUG(dbgs() << "\nAdd dep type " << getDependenceTypeStr(dt) << ":\t ("
119                     << *indexingOpView.op << ", " << indexingOpView.operandIndex
120                     << ") -> \n\t\t(" << *dependentOpView.op << ", "
121                     << dependentOpView.operandIndex << ")");
122   dependencesFromGraphs[dt][indexingOpView.op].push_back(
123       LinalgDependenceGraphElem{dependentOpView, indexingOpView, dt});
124   dependencesIntoGraphs[dt][dependentOpView.op].push_back(
125       LinalgDependenceGraphElem{indexingOpView, dependentOpView, dt});
126 }
127 
128 LinalgDependenceGraph::dependence_range
129 LinalgDependenceGraph::getDependencesFrom(
130     LinalgOp src, LinalgDependenceGraph::DependenceType dt) const {
131   return getDependencesFrom(src.getOperation(), dt);
132 }
133 
134 LinalgDependenceGraph::dependence_range
135 LinalgDependenceGraph::getDependencesFrom(
136     Operation *src, LinalgDependenceGraph::DependenceType dt) const {
137   auto iter = dependencesFromGraphs[dt].find(src);
138   if (iter == dependencesFromGraphs[dt].end())
139     return llvm::make_range(nullptr, nullptr);
140   return llvm::make_range(iter->second.begin(), iter->second.end());
141 }
142 
143 LinalgDependenceGraph::dependence_range
144 LinalgDependenceGraph::getDependencesInto(
145     LinalgOp dst, LinalgDependenceGraph::DependenceType dt) const {
146   return getDependencesInto(dst.getOperation(), dt);
147 }
148 
149 LinalgDependenceGraph::dependence_range
150 LinalgDependenceGraph::getDependencesInto(
151     Operation *dst, LinalgDependenceGraph::DependenceType dt) const {
152   auto iter = dependencesIntoGraphs[dt].find(dst);
153   if (iter == dependencesIntoGraphs[dt].end())
154     return llvm::make_range(nullptr, nullptr);
155   return llvm::make_range(iter->second.begin(), iter->second.end());
156 }
157 
158 void LinalgDependenceGraph::addDependencesBetween(LinalgOp src, LinalgOp dst) {
159   for (auto srcView : llvm::enumerate(src.getOutputBuffers())) { // W
160     unsigned srcIndex =
161         src.getOperandIndexForOutputIndex(srcView.index()).getValue();
162     // RAW graph
163     for (auto dstView : llvm::enumerate(dst.getInputBuffers())) { // R
164       if (aliases.alias(srcView.value(),
165                         dstView.value())) { // if alias, fill RAW
166         unsigned dstIndex =
167             dst.getOperandIndexForInputIndex(dstView.index()).getValue();
168         addDependenceElem(DependenceType::RAW,
169                           LinalgOpView{src.getOperation(), srcIndex},
170                           LinalgOpView{dst.getOperation(), dstIndex});
171       }
172     }
173     // WAW graph
174     for (auto dstView : llvm::enumerate(dst.getOutputBuffers())) { // W
175       if (aliases.alias(srcView.value(),
176                         dstView.value())) { // if alias, fill WAW
177         unsigned dstIndex =
178             dst.getOperandIndexForOutputIndex(dstView.index()).getValue();
179         addDependenceElem(DependenceType::WAW,
180                           LinalgOpView{src.getOperation(), srcIndex},
181                           LinalgOpView{dst.getOperation(), dstIndex});
182       }
183     }
184   }
185   for (auto srcView : llvm::enumerate(src.getInputBuffers())) { // R
186     unsigned srcIndex =
187         src.getOperandIndexForInputIndex(srcView.index()).getValue();
188     // RAR graph
189     for (auto dstView : llvm::enumerate(dst.getInputBuffers())) { // R
190       if (aliases.alias(srcView.value(),
191                         dstView.value())) { // if alias, fill RAR
192         unsigned dstIndex =
193             dst.getOperandIndexForInputIndex(dstView.index()).getValue();
194         addDependenceElem(DependenceType::RAR,
195                           LinalgOpView{src.getOperation(), srcIndex},
196                           LinalgOpView{dst.getOperation(), dstIndex});
197       }
198     }
199     // WAR graph
200     for (auto dstView : llvm::enumerate(dst.getOutputBuffers())) { // W
201       if (aliases.alias(srcView.value(),
202                         dstView.value())) { // if alias, fill WAR
203         unsigned dstIndex =
204             dst.getOperandIndexForOutputIndex(dstView.index()).getValue();
205         addDependenceElem(DependenceType::WAR,
206                           LinalgOpView{src.getOperation(), srcIndex},
207                           LinalgOpView{dst.getOperation(), dstIndex});
208       }
209     }
210   }
211 }
212 
213 SmallVector<Operation *, 8>
214 LinalgDependenceGraph::findCoveringDependences(LinalgOp srcLinalgOp,
215                                                LinalgOp dstLinalgOp) const {
216   return findOperationsWithCoveringDependences(
217       srcLinalgOp, dstLinalgOp, nullptr,
218       {DependenceType::WAW, DependenceType::WAR, DependenceType::RAW});
219 }
220 
221 SmallVector<Operation *, 8> LinalgDependenceGraph::findCoveringWrites(
222     LinalgOp srcLinalgOp, LinalgOp dstLinalgOp, Value view) const {
223   return findOperationsWithCoveringDependences(
224       srcLinalgOp, dstLinalgOp, view,
225       {DependenceType::WAW, DependenceType::WAR});
226 }
227 
228 SmallVector<Operation *, 8> LinalgDependenceGraph::findCoveringReads(
229     LinalgOp srcLinalgOp, LinalgOp dstLinalgOp, Value view) const {
230   return findOperationsWithCoveringDependences(
231       srcLinalgOp, dstLinalgOp, view,
232       {DependenceType::RAR, DependenceType::RAW});
233 }
234 
235 SmallVector<Operation *, 8>
236 LinalgDependenceGraph::findOperationsWithCoveringDependences(
237     LinalgOp srcLinalgOp, LinalgOp dstLinalgOp, Value view,
238     ArrayRef<DependenceType> types) const {
239   auto *src = srcLinalgOp.getOperation();
240   auto *dst = dstLinalgOp.getOperation();
241   auto srcPos = linalgOpPositions.lookup(src);
242   auto dstPos = linalgOpPositions.lookup(dst);
243   assert(srcPos < dstPos && "expected dst after src in IR traversal order");
244 
245   SmallVector<Operation *, 8> res;
246   // Consider an intermediate interleaved `interim` op, look for any dependence
247   // to an aliasing view on a src -> op -> dst path.
248   // TODO: we are not considering paths yet, just interleaved positions.
249   for (auto dt : types) {
250     for (auto dependence : getDependencesFrom(src, dt)) {
251       auto interimPos = linalgOpPositions.lookup(dependence.dependentOpView.op);
252       // Skip if not interleaved.
253       if (interimPos >= dstPos || interimPos <= srcPos)
254         continue;
255       linalg::LinalgOp consumer =
256           cast<linalg::LinalgOp>(dependence.indexingOpView.op);
257       Value consumerView =
258           consumer.getShapedOperand(dependence.indexingOpView.operandIndex);
259       if (view && !aliases.alias(view, consumerView))
260         continue;
261       auto *op = dependence.dependentOpView.op;
262       LLVM_DEBUG(dbgs() << "\n***Found covering dependence of type "
263                         << getDependenceTypeStr(dt) << ": " << *src << " -> "
264                         << *op << " on " << consumerView);
265       res.push_back(op);
266     }
267   }
268   return res;
269 }
270 
271 bool LinalgDependenceGraph::hasDependenceFrom(
272     LinalgOp srcLinalgOp, LinalgOp dstLinalgOp,
273     ArrayRef<LinalgDependenceGraph::DependenceType> depTypes) const {
274   for (auto dep : depTypes) {
275     for (auto dependence : getDependencesInto(dstLinalgOp, dep)) {
276       if (dependence.dependentOpView.op == srcLinalgOp)
277         return true;
278     }
279   }
280   return false;
281 }
282 
283 bool LinalgDependenceGraph::hasDependentOperationsFrom(
284     LinalgOp linalgOp,
285     ArrayRef<LinalgDependenceGraph::DependenceType> depTypes) const {
286   for (auto dep : depTypes) {
287     if (!getDependencesFrom(linalgOp, dep).empty())
288       return true;
289   }
290   return false;
291 }
292 
293 bool LinalgDependenceGraph::hasDependentOperationsInto(
294     LinalgOp linalgOp,
295     ArrayRef<LinalgDependenceGraph::DependenceType> depTypes) const {
296   for (auto dep : depTypes) {
297     if (!getDependencesInto(linalgOp, dep).empty())
298       return true;
299   }
300   return false;
301 }
302 
303 bool LinalgDependenceGraph::hasDependentOperations(
304     LinalgOp linalgOp, ArrayRef<DependenceType> depTypes) const {
305   return hasDependentOperationsInto(linalgOp, depTypes) ||
306          hasDependentOperationsFrom(linalgOp, depTypes);
307 }
308 
309 SmallVector<LinalgDependenceGraph::LinalgDependenceGraphElem, 2>
310 LinalgDependenceGraph::getDependentOperationsInto(
311     LinalgOp linalgOp, ArrayRef<DependenceType> depTypes) const {
312   SmallVector<LinalgDependenceGraph::LinalgDependenceGraphElem, 2>
313       dependentOperations;
314   for (auto dependenceType : depTypes) {
315     auto dependencies = getDependencesInto(linalgOp, dependenceType);
316     dependentOperations.append(dependencies.begin(), dependencies.end());
317   }
318   return dependentOperations;
319 }
320 
321 SmallVector<LinalgDependenceGraph::LinalgDependenceGraphElem, 2>
322 LinalgDependenceGraph::getDependentOperationsFrom(
323     LinalgOp linalgOp, ArrayRef<DependenceType> depTypes) const {
324   SmallVector<LinalgDependenceGraph::LinalgDependenceGraphElem, 2>
325       dependentOperations;
326   for (auto dependenceType : depTypes) {
327     auto dependencies = getDependencesFrom(linalgOp, dependenceType);
328     dependentOperations.append(dependencies.begin(), dependencies.end());
329   }
330   return dependentOperations;
331 }
332 
333 /// Returns all dependent operations (into and from) given `operation`.
334 SmallVector<LinalgDependenceGraph::LinalgDependenceGraphElem, 2>
335 LinalgDependenceGraph::getDependentOperations(
336     LinalgOp linalgOp, ArrayRef<DependenceType> depTypes) const {
337   SmallVector<LinalgDependenceGraphElem, 2> dependentOperations =
338       getDependentOperationsInto(linalgOp, depTypes);
339   SmallVector<LinalgDependenceGraphElem, 2> t =
340       getDependentOperationsFrom(linalgOp, depTypes);
341   dependentOperations.append(t.begin(), t.end());
342   return dependentOperations;
343 }
344