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     if (auto memEffect = dyn_cast<MemoryEffectOpInterface>(defOp)) {
47       // Collect all memory effects on `v`.
48       SmallVector<MemoryEffects::EffectInstance, 1> effects;
49       memEffect.getEffectsOnValue(v, effects);
50 
51       // If we have the 'Allocate' memory effect on `v`, then `v` should be the
52       // original buffer.
53       if (llvm::any_of(
54               effects, [](const MemoryEffects::EffectInstance &instance) {
55                 return isa<MemoryEffects::Allocate>(instance.getEffect());
56               }))
57         return v;
58     }
59 
60     if (auto viewLikeOp = dyn_cast<ViewLikeOpInterface>(defOp)) {
61       auto it =
62           aliases.insert(std::make_pair(v, find(viewLikeOp.getViewSource())));
63       return it.first->second;
64     }
65 
66     llvm::errs() << "View alias analysis reduces to: " << v << "\n";
67     llvm_unreachable("unsupported view alias case");
68   }
69 }
70 
71 StringRef LinalgDependenceGraph::getDependenceTypeStr(DependenceType depType) {
72   switch (depType) {
73   case LinalgDependenceGraph::DependenceType::RAW:
74     return "RAW";
75   case LinalgDependenceGraph::DependenceType::RAR:
76     return "RAR";
77   case LinalgDependenceGraph::DependenceType::WAR:
78     return "WAR";
79   case LinalgDependenceGraph::DependenceType::WAW:
80     return "WAW";
81   default:
82     break;
83   }
84   llvm_unreachable("Unexpected DependenceType");
85 }
86 
87 LinalgDependenceGraph
88 LinalgDependenceGraph::buildDependenceGraph(Aliases &aliases, FuncOp f) {
89   SmallVector<Operation *, 8> linalgOps;
90   f.walk([&](LinalgOp op) { linalgOps.push_back(op); });
91   return LinalgDependenceGraph(aliases, linalgOps);
92 }
93 
94 LinalgDependenceGraph::LinalgDependenceGraph(Aliases &aliases,
95                                              ArrayRef<Operation *> ops)
96     : aliases(aliases), linalgOps(ops.begin(), ops.end()) {
97   for (auto en : llvm::enumerate(linalgOps)) {
98     assert(isa<LinalgOp>(en.value()) && "Expected value for LinalgOp");
99     linalgOpPositions.insert(std::make_pair(en.value(), en.index()));
100   }
101   for (unsigned i = 0, e = ops.size(); i < e; ++i) {
102     for (unsigned j = i + 1; j < e; ++j) {
103       addDependencesBetween(cast<LinalgOp>(ops[i]), cast<LinalgOp>(ops[j]));
104     }
105   }
106 }
107 
108 void LinalgDependenceGraph::addDependenceElem(DependenceType dt,
109                                               LinalgOpView indexingOpView,
110                                               LinalgOpView dependentOpView) {
111   LLVM_DEBUG(dbgs() << "\nAdd dep type " << getDependenceTypeStr(dt) << ":\t"
112                     << *indexingOpView.op << " -> " << *dependentOpView.op);
113   dependencesFromGraphs[dt][indexingOpView.op].push_back(
114       LinalgDependenceGraphElem{dependentOpView, indexingOpView.view});
115   dependencesIntoGraphs[dt][dependentOpView.op].push_back(
116       LinalgDependenceGraphElem{indexingOpView, dependentOpView.view});
117 }
118 
119 LinalgDependenceGraph::dependence_range
120 LinalgDependenceGraph::getDependencesFrom(
121     LinalgOp src, LinalgDependenceGraph::DependenceType dt) const {
122   return getDependencesFrom(src.getOperation(), dt);
123 }
124 
125 LinalgDependenceGraph::dependence_range
126 LinalgDependenceGraph::getDependencesFrom(
127     Operation *src, LinalgDependenceGraph::DependenceType dt) const {
128   auto iter = dependencesFromGraphs[dt].find(src);
129   if (iter == dependencesFromGraphs[dt].end())
130     return llvm::make_range(nullptr, nullptr);
131   return llvm::make_range(iter->second.begin(), iter->second.end());
132 }
133 
134 LinalgDependenceGraph::dependence_range
135 LinalgDependenceGraph::getDependencesInto(
136     LinalgOp dst, LinalgDependenceGraph::DependenceType dt) const {
137   return getDependencesInto(dst.getOperation(), dt);
138 }
139 
140 LinalgDependenceGraph::dependence_range
141 LinalgDependenceGraph::getDependencesInto(
142     Operation *dst, LinalgDependenceGraph::DependenceType dt) const {
143   auto iter = dependencesIntoGraphs[dt].find(dst);
144   if (iter == dependencesIntoGraphs[dt].end())
145     return llvm::make_range(nullptr, nullptr);
146   return llvm::make_range(iter->second.begin(), iter->second.end());
147 }
148 
149 void LinalgDependenceGraph::addDependencesBetween(LinalgOp src, LinalgOp dst) {
150   for (auto srcView : src.getOutputBuffers()) { // W
151     // RAW graph
152     for (auto dstView : dst.getInputBuffers()) { // R
153       if (aliases.alias(srcView, dstView)) { // if alias, fill RAW
154         addDependenceElem(DependenceType::RAW,
155                           LinalgOpView{src.getOperation(), srcView},
156                           LinalgOpView{dst.getOperation(), dstView});
157       }
158     }
159     // WAW graph
160     for (auto dstView : dst.getOutputBuffers()) { // W
161       if (aliases.alias(srcView, dstView)) {      // if alias, fill WAW
162         addDependenceElem(DependenceType::WAW,
163                           LinalgOpView{src.getOperation(), srcView},
164                           LinalgOpView{dst.getOperation(), dstView});
165       }
166     }
167   }
168   for (auto srcView : src.getInputBuffers()) { // R
169     // RAR graph
170     for (auto dstView : dst.getInputBuffers()) { // R
171       if (aliases.alias(srcView, dstView)) { // if alias, fill RAR
172         addDependenceElem(DependenceType::RAR,
173                           LinalgOpView{src.getOperation(), srcView},
174                           LinalgOpView{dst.getOperation(), dstView});
175       }
176     }
177     // WAR graph
178     for (auto dstView : dst.getOutputBuffers()) { // W
179       if (aliases.alias(srcView, dstView)) {      // if alias, fill WAR
180         addDependenceElem(DependenceType::WAR,
181                           LinalgOpView{src.getOperation(), srcView},
182                           LinalgOpView{dst.getOperation(), dstView});
183       }
184     }
185   }
186 }
187 
188 SmallVector<Operation *, 8>
189 LinalgDependenceGraph::findCoveringDependences(LinalgOp srcLinalgOp,
190                                                LinalgOp dstLinalgOp) const {
191   return findOperationsWithCoveringDependences(
192       srcLinalgOp, dstLinalgOp, nullptr,
193       {DependenceType::WAW, DependenceType::WAR, DependenceType::RAW});
194 }
195 
196 SmallVector<Operation *, 8> LinalgDependenceGraph::findCoveringWrites(
197     LinalgOp srcLinalgOp, LinalgOp dstLinalgOp, Value view) const {
198   return findOperationsWithCoveringDependences(
199       srcLinalgOp, dstLinalgOp, view,
200       {DependenceType::WAW, DependenceType::WAR});
201 }
202 
203 SmallVector<Operation *, 8> LinalgDependenceGraph::findCoveringReads(
204     LinalgOp srcLinalgOp, LinalgOp dstLinalgOp, Value view) const {
205   return findOperationsWithCoveringDependences(
206       srcLinalgOp, dstLinalgOp, view,
207       {DependenceType::RAR, DependenceType::RAW});
208 }
209 
210 SmallVector<Operation *, 8>
211 LinalgDependenceGraph::findOperationsWithCoveringDependences(
212     LinalgOp srcLinalgOp, LinalgOp dstLinalgOp, Value view,
213     ArrayRef<DependenceType> types) const {
214   auto *src = srcLinalgOp.getOperation();
215   auto *dst = dstLinalgOp.getOperation();
216   auto srcPos = linalgOpPositions.lookup(src);
217   auto dstPos = linalgOpPositions.lookup(dst);
218   assert(srcPos < dstPos && "expected dst after src in IR traversal order");
219 
220   SmallVector<Operation *, 8> res;
221   // Consider an intermediate interleaved `interim` op, look for any dependence
222   // to an aliasing view on a src -> op -> dst path.
223   // TODO: we are not considering paths yet, just interleaved positions.
224   for (auto dt : types) {
225     for (auto dependence : getDependencesFrom(src, dt)) {
226       auto interimPos = linalgOpPositions.lookup(dependence.dependentOpView.op);
227       // Skip if not interleaved.
228       if (interimPos >= dstPos || interimPos <= srcPos)
229         continue;
230       if (view && !aliases.alias(view, dependence.indexingView))
231         continue;
232       auto *op = dependence.dependentOpView.op;
233       LLVM_DEBUG(dbgs() << "\n***Found covering dependence of type "
234                         << getDependenceTypeStr(dt) << ": " << *src << " -> "
235                         << *op << " on " << dependence.indexingView);
236       res.push_back(op);
237     }
238   }
239   return res;
240 }
241 
242 bool LinalgDependenceGraph::hasDependenceFrom(
243     LinalgOp srcLinalgOp, LinalgOp dstLinalgOp,
244     ArrayRef<LinalgDependenceGraph::DependenceType> depTypes) const {
245   for (auto dep : depTypes) {
246     for (auto dependence : getDependencesInto(dstLinalgOp, dep)) {
247       if (dependence.dependentOpView.op == srcLinalgOp)
248         return true;
249     }
250   }
251   return false;
252 }
253 
254 bool LinalgDependenceGraph::hasDependentOperations(
255     LinalgOp linalgOp,
256     ArrayRef<LinalgDependenceGraph::DependenceType> depTypes) const {
257   for (auto dep : depTypes) {
258     if (!getDependencesFrom(linalgOp, dep).empty() ||
259         !getDependencesInto(linalgOp, dep).empty())
260       return true;
261   }
262   return false;
263 }
264