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