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