MLIR 24.0.0git
Utils.cpp
Go to the documentation of this file.
1//===- Utils.cpp ---- Misc utilities for analysis -------------------------===//
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 miscellaneous analysis routines for non-loop IR
10// structures.
11//
12//===----------------------------------------------------------------------===//
13
15
23#include "mlir/IR/IntegerSet.h"
24#include "llvm/ADT/SetVector.h"
25#include "llvm/ADT/SmallVectorExtras.h"
26#include "llvm/Support/Debug.h"
27#include "llvm/Support/DebugLog.h"
28#include "llvm/Support/raw_ostream.h"
29#include <optional>
30
31#define DEBUG_TYPE "analysis-utils"
32
33using namespace mlir;
34using namespace affine;
35using namespace presburger;
36
37using llvm::SmallDenseMap;
38
40
41// LoopNestStateCollector walks loop nests and collects load and store
42// operations, and whether or not a region holding op other than ForOp and IfOp
43// was encountered in the loop nest.
45 opToWalk->walk([&](Operation *op) {
46 if (auto forOp = dyn_cast<AffineForOp>(op)) {
47 forOps.push_back(forOp);
48 } else if (isa<AffineReadOpInterface>(op)) {
49 loadOpInsts.push_back(op);
50 } else if (isa<AffineWriteOpInterface>(op)) {
51 storeOpInsts.push_back(op);
52 } else {
53 auto memInterface = dyn_cast<MemoryEffectOpInterface>(op);
54 if (!memInterface) {
56 // This op itself is memory-effect free.
57 return;
58 // Check operands. Eg. ops like the `call` op are handled here.
59 for (Value v : op->getOperands()) {
60 if (!isa<MemRefType>(v.getType()))
61 continue;
62 // Conservatively, we assume the memref is read and written to.
63 memrefLoads.push_back(op);
64 memrefStores.push_back(op);
65 }
66 } else {
67 // Non-affine loads and stores.
69 memrefLoads.push_back(op);
71 memrefStores.push_back(op);
73 memrefFrees.push_back(op);
74 }
75 }
76 });
77}
78
80 unsigned loadOpCount = 0;
81 for (Operation *loadOp : loads) {
82 // Common case: affine reads.
83 if (auto affineLoad = dyn_cast<AffineReadOpInterface>(loadOp)) {
84 if (memref == affineLoad.getMemRef())
85 ++loadOpCount;
86 } else if (hasEffect<MemoryEffects::Read>(loadOp, memref)) {
87 ++loadOpCount;
88 }
89 }
90 return loadOpCount;
91}
92
93// Returns the store op count for 'memref'.
95 unsigned storeOpCount = 0;
96 for (auto *storeOp : llvm::concat<Operation *const>(stores, memrefStores)) {
97 // Common case: affine writes.
98 if (auto affineStore = dyn_cast<AffineWriteOpInterface>(storeOp)) {
99 if (memref == affineStore.getMemRef())
100 ++storeOpCount;
101 } else if (hasEffect<MemoryEffects::Write>(const_cast<Operation *>(storeOp),
102 memref)) {
103 ++storeOpCount;
104 }
105 }
106 return storeOpCount;
107}
108
109// Returns the store op count for 'memref'.
110unsigned Node::hasStore(Value memref) const {
111 return llvm::any_of(
112 llvm::concat<Operation *const>(stores, memrefStores),
113 [&](Operation *storeOp) {
114 if (auto affineStore = dyn_cast<AffineWriteOpInterface>(storeOp)) {
115 if (memref == affineStore.getMemRef())
116 return true;
117 } else if (hasEffect<MemoryEffects::Write>(storeOp, memref)) {
118 return true;
119 }
120 return false;
121 });
122}
123
124unsigned Node::hasFree(Value memref) const {
125 return llvm::any_of(memrefFrees, [&](Operation *freeOp) {
127 });
128}
129
130// Returns all store ops in 'storeOps' which access 'memref'.
132 SmallVectorImpl<Operation *> *storeOps) const {
133 for (Operation *storeOp : stores) {
134 if (memref == cast<AffineWriteOpInterface>(storeOp).getMemRef())
135 storeOps->push_back(storeOp);
136 }
137}
138
139// Returns all load ops in 'loadOps' which access 'memref'.
141 SmallVectorImpl<Operation *> *loadOps) const {
142 for (Operation *loadOp : loads) {
143 if (memref == cast<AffineReadOpInterface>(loadOp).getMemRef())
144 loadOps->push_back(loadOp);
145 }
146}
147
148// Returns all memrefs in 'loadAndStoreMemrefSet' for which this node
149// has at least one load and store operation.
151 DenseSet<Value> *loadAndStoreMemrefSet) const {
152 llvm::SmallDenseSet<Value, 2> loadMemrefs;
153 for (Operation *loadOp : loads) {
154 loadMemrefs.insert(cast<AffineReadOpInterface>(loadOp).getMemRef());
155 }
156 for (Operation *storeOp : stores) {
157 auto memref = cast<AffineWriteOpInterface>(storeOp).getMemRef();
158 if (loadMemrefs.count(memref) > 0)
159 loadAndStoreMemrefSet->insert(memref);
160 }
161}
162
163/// Returns the values that this op has a memref effect of type `EffectTys` on,
164/// not considering recursive effects.
165template <typename... EffectTys>
167 auto memOp = dyn_cast<MemoryEffectOpInterface>(op);
168 if (!memOp) {
170 // No effects.
171 return;
172 // Memref operands have to be considered as being affected.
173 for (Value operand : op->getOperands()) {
174 if (isa<MemRefType>(operand.getType()))
175 values.push_back(operand);
176 }
177 return;
178 }
180 memOp.getEffects(effects);
181 for (auto &effect : effects) {
182 Value effectVal = effect.getValue();
183 if (isa<EffectTys...>(effect.getEffect()) && effectVal &&
184 isa<MemRefType>(effectVal.getType()))
185 values.push_back(effectVal);
186 };
187}
188
189/// Add `op` to MDG creating a new node and adding its memory accesses (affine
190/// or non-affine to memrefAccesses (memref -> list of nodes with accesses) map.
191static Node *
193 DenseMap<Value, SetVector<unsigned>> &memrefAccesses) {
194 auto &nodes = mdg.nodes;
195 // Create graph node 'id' to represent top-level 'forOp' and record
196 // all loads and store accesses it contains.
197 LoopNestStateCollector collector;
198 collector.collect(nodeOp);
199 unsigned newNodeId = mdg.nextNodeId++;
200 Node &node = nodes.insert({newNodeId, Node(newNodeId, nodeOp)}).first->second;
201 for (Operation *op : collector.loadOpInsts) {
202 node.loads.push_back(op);
203 auto memref = cast<AffineReadOpInterface>(op).getMemRef();
204 memrefAccesses[memref].insert(node.id);
205 }
206 for (Operation *op : collector.storeOpInsts) {
207 node.stores.push_back(op);
208 auto memref = cast<AffineWriteOpInterface>(op).getMemRef();
209 memrefAccesses[memref].insert(node.id);
210 }
211 for (Operation *op : collector.memrefLoads) {
212 SmallVector<Value> effectedValues;
213 getEffectedValues<MemoryEffects::Read>(op, effectedValues);
214 if (llvm::any_of(((ValueRange)effectedValues).getTypes(),
215 [](Type type) { return !isa<MemRefType>(type); }))
216 // We do not know the interaction here.
217 return nullptr;
218 for (Value memref : effectedValues)
219 memrefAccesses[memref].insert(node.id);
220 node.memrefLoads.push_back(op);
221 }
222 for (Operation *op : collector.memrefStores) {
223 SmallVector<Value> effectedValues;
224 getEffectedValues<MemoryEffects::Write>(op, effectedValues);
225 if (llvm::any_of((ValueRange(effectedValues)).getTypes(),
226 [](Type type) { return !isa<MemRefType>(type); }))
227 return nullptr;
228 for (Value memref : effectedValues)
229 memrefAccesses[memref].insert(node.id);
230 node.memrefStores.push_back(op);
231 }
232 for (Operation *op : collector.memrefFrees) {
233 SmallVector<Value> effectedValues;
234 getEffectedValues<MemoryEffects::Free>(op, effectedValues);
235 if (llvm::any_of((ValueRange(effectedValues)).getTypes(),
236 [](Type type) { return !isa<MemRefType>(type); }))
237 return nullptr;
238 for (Value memref : effectedValues)
239 memrefAccesses[memref].insert(node.id);
240 node.memrefFrees.push_back(op);
241 }
242
243 return &node;
244}
245
246/// Returns true if `op` may read from or write to `memref`.
248 SmallVector<Value> effectedValues;
250 effectedValues);
251 if (llvm::is_contained(effectedValues, memref))
252 return true;
253
254 auto memoryEffectOp = dyn_cast<MemoryEffectOpInterface>(op);
255 if (!memoryEffectOp)
256 return false;
257
259 memoryEffectOp.getEffects(effects);
260 return llvm::any_of(effects, [](const auto &effect) {
261 if (!isa<MemoryEffects::Read, MemoryEffects::Write>(effect.getEffect()))
262 return false;
263
264 // Value-less read/write effects may target unknown memory. Since this is a
265 // may-analysis, conservatively assume they may access the memref.
266 return !effect.getValue();
267 });
268}
269
270/// Returns true if there may be a dependence on `memref` from srcNode's
271/// memory ops to dstNode's memory ops, while using the affine memory
272/// dependence analysis checks. The method assumes that there is at least one
273/// memory op in srcNode's loads and stores on `memref`, and similarly for
274/// `dstNode`. `srcNode.op` and `destNode.op` are expected to be nested in the
275/// same block and so the dependences are tested at the depth of that block.
276static bool mayDependence(const Node &srcNode, const Node &dstNode,
277 Value memref) {
278 assert(srcNode.op->getBlock() == dstNode.op->getBlock());
279 if (!isa<AffineForOp>(srcNode.op) || !isa<AffineForOp>(dstNode.op))
280 return true;
281
282 // Conservatively handle dependences involving non-affine load/stores. Return
283 // true if there exists a conflicting read/write access involving such.
284
285 // Check whether there is a dependence from a source read/write op to a
286 // destination read/write op on `memref`.
287 auto hasNonAffineDep = [&](ArrayRef<Operation *> srcMemOps,
288 ArrayRef<Operation *> dstMemOps) {
289 return llvm::any_of(srcMemOps,
290 [&](Operation *srcOp) {
291 return mayAccessMemRef(srcOp, memref);
292 }) &&
293 llvm::any_of(dstMemOps, [&](Operation *dstOp) {
294 return mayAccessMemRef(dstOp, memref);
295 });
296 };
297
299 // Between non-affine src stores and dst load/store.
300 llvm::append_range(dstOps, llvm::concat<Operation *const>(
301 dstNode.loads, dstNode.stores,
302 dstNode.memrefLoads, dstNode.memrefStores));
303 if (hasNonAffineDep(srcNode.memrefStores, dstOps))
304 return true;
305 // Between non-affine loads and dst stores.
306 dstOps.clear();
307 llvm::append_range(dstOps, llvm::concat<Operation *const>(
308 dstNode.stores, dstNode.memrefStores));
309 if (hasNonAffineDep(srcNode.memrefLoads, dstOps))
310 return true;
311 // Between affine stores and memref load/stores.
312 dstOps.clear();
313 llvm::append_range(dstOps, llvm::concat<Operation *const>(
314 dstNode.memrefLoads, dstNode.memrefStores));
315 if (hasNonAffineDep(srcNode.stores, dstOps))
316 return true;
317 // Between affine loads and memref stores.
318 dstOps.clear();
319 llvm::append_range(dstOps, dstNode.memrefStores);
320 if (hasNonAffineDep(srcNode.loads, dstOps))
321 return true;
322
323 // Affine load/store pairs. We don't need to check for locally allocated
324 // memrefs since the dependence analysis here is between mem ops from
325 // srcNode's for op to dstNode's for op at the depth at which those
326 // `affine.for` ops are nested, i.e., dependences at depth `d + 1` where
327 // `d` is the number of common surrounding loops.
328 for (auto *srcMemOp :
329 llvm::concat<Operation *const>(srcNode.stores, srcNode.loads)) {
330 MemRefAccess srcAcc(srcMemOp);
331 if (srcAcc.memref != memref)
332 continue;
333 for (auto *destMemOp :
334 llvm::concat<Operation *const>(dstNode.stores, dstNode.loads)) {
335 MemRefAccess destAcc(destMemOp);
336 if (destAcc.memref != memref)
337 continue;
338 // Check for a top-level dependence between srcNode and destNode's ops.
340 srcAcc, destAcc, getNestingDepth(srcNode.op) + 1)))
341 return true;
342 }
343 }
344 return false;
345}
346
347bool MemRefDependenceGraph::init(bool fullAffineDependences) {
348 LDBG() << "--- Initializing MDG ---";
349 // Map from a memref to the set of ids of the nodes that have ops accessing
350 // the memref.
352
353 // Create graph nodes.
355 for (Operation &op : block) {
356 if (auto forOp = dyn_cast<AffineForOp>(op)) {
357 Node *node = addNodeToMDG(&op, *this, memrefAccesses);
358 if (!node)
359 return false;
360 forToNodeMap[&op] = node->id;
361 } else if (isa<AffineReadOpInterface>(op)) {
362 // Create graph node for top-level load op.
363 Node node(nextNodeId++, &op);
364 node.loads.push_back(&op);
365 auto memref = cast<AffineReadOpInterface>(op).getMemRef();
366 memrefAccesses[memref].insert(node.id);
367 nodes.insert({node.id, node});
368 } else if (isa<AffineWriteOpInterface>(op)) {
369 // Create graph node for top-level store op.
370 Node node(nextNodeId++, &op);
371 node.stores.push_back(&op);
372 auto memref = cast<AffineWriteOpInterface>(op).getMemRef();
373 memrefAccesses[memref].insert(node.id);
374 nodes.insert({node.id, node});
375 } else if (op.getNumResults() > 0 && !op.use_empty()) {
376 // Create graph node for top-level producer of SSA values, which
377 // could be used by loop nest nodes.
378 Node *node = addNodeToMDG(&op, *this, memrefAccesses);
379 if (!node)
380 return false;
381 } else if (!isMemoryEffectFree(&op) &&
382 (op.getNumRegions() == 0 || isa<RegionBranchOpInterface>(op))) {
383 // Create graph node for top-level op unless it is known to be
384 // memory-effect free. This covers all unknown/unregistered ops,
385 // non-affine ops with memory effects, and region-holding ops with a
386 // well-defined control flow. During the fusion validity checks, edges
387 // to/from these ops get looked at.
388 Node *node = addNodeToMDG(&op, *this, memrefAccesses);
389 if (!node)
390 return false;
391 } else if (op.getNumRegions() != 0 && !isa<RegionBranchOpInterface>(op)) {
392 // Return false if non-handled/unknown region-holding ops are found. We
393 // won't know what such ops do or what its regions mean; for e.g., it may
394 // not be an imperative op.
395 LDBG() << "MDG init failed; unknown region-holding op found!";
396 return false;
397 }
398 // We aren't creating nodes for memory-effect free ops either with no
399 // regions (unless it has results being used) or those with branch op
400 // interface.
401 }
402
403 LDBG() << "Created " << nodes.size() << " nodes";
404
405 // Add dependence edges between nodes which produce SSA values and their
406 // users. Load ops can be considered as the ones producing SSA values.
407 for (auto &idAndNode : nodes) {
408 const Node &node = idAndNode.second;
409 // Stores don't define SSA values, skip them.
410 if (!node.stores.empty())
411 continue;
412 Operation *opInst = node.op;
413 for (Value value : opInst->getResults()) {
414 for (Operation *user : value.getUsers()) {
415 // Ignore users outside of the block.
416 if (block.getParent()->findAncestorOpInRegion(*user)->getBlock() !=
417 &block)
418 continue;
420 getAffineForIVs(*user, &loops);
421 // Find the surrounding affine.for nested immediately within the
422 // block.
423 auto *it = llvm::find_if(loops, [&](AffineForOp loop) {
424 return loop->getBlock() == &block;
425 });
426 if (it == loops.end())
427 continue;
428 assert(forToNodeMap.count(*it) > 0 && "missing mapping");
429 unsigned userLoopNestId = forToNodeMap[*it];
430 addEdge(node.id, userLoopNestId, value);
431 }
432 }
433 }
434
435 // Walk memref access lists and add graph edges between dependent nodes.
436 for (auto &memrefAndList : memrefAccesses) {
437 unsigned n = memrefAndList.second.size();
438 Value srcMemRef = memrefAndList.first;
439 // Add edges between all dependent pairs among the node IDs on this memref.
440 for (unsigned i = 0; i < n; ++i) {
441 unsigned srcId = memrefAndList.second[i];
442 Node *srcNode = getNode(srcId);
443 bool srcHasStoreOrFree =
444 srcNode->hasStore(srcMemRef) || srcNode->hasFree(srcMemRef);
445 for (unsigned j = i + 1; j < n; ++j) {
446 unsigned dstId = memrefAndList.second[j];
447 Node *dstNode = getNode(dstId);
448 bool dstHasStoreOrFree =
449 dstNode->hasStore(srcMemRef) || dstNode->hasFree(srcMemRef);
450 if ((srcHasStoreOrFree || dstHasStoreOrFree)) {
451 // Check precise affine deps if asked for; otherwise, conservative.
452 if (!fullAffineDependences ||
453 mayDependence(*srcNode, *dstNode, srcMemRef))
454 addEdge(srcId, dstId, srcMemRef);
455 }
456 }
457 }
458 }
459 return true;
460}
461
462// Returns the graph node for 'id'.
463const Node *MemRefDependenceGraph::getNode(unsigned id) const {
464 auto it = nodes.find(id);
465 assert(it != nodes.end());
466 return &it->second;
467}
468
469// Returns the graph node for 'forOp'.
470const Node *MemRefDependenceGraph::getForOpNode(AffineForOp forOp) const {
471 for (auto &idAndNode : nodes)
472 if (idAndNode.second.op == forOp)
473 return &idAndNode.second;
474 return nullptr;
475}
476
477// Adds a node with 'op' to the graph and returns its unique identifier.
479 Node node(nextNodeId++, op);
480 nodes.insert({node.id, node});
481 return node.id;
482}
483
484// Remove node 'id' (and its associated edges) from graph.
486 // Remove each edge in 'inEdges[id]'.
487 if (inEdges.count(id) > 0) {
488 SmallVector<Edge, 2> oldInEdges = inEdges[id];
489 for (auto &inEdge : oldInEdges) {
490 removeEdge(inEdge.id, id, inEdge.value);
491 }
492 }
493 // Remove each edge in 'outEdges[id]'.
494 if (outEdges.contains(id)) {
495 SmallVector<Edge, 2> oldOutEdges = outEdges[id];
496 for (auto &outEdge : oldOutEdges) {
497 removeEdge(id, outEdge.id, outEdge.value);
498 }
499 }
500 // Erase remaining node state.
501 inEdges.erase(id);
502 outEdges.erase(id);
503 nodes.erase(id);
504}
505
506// Returns true if node 'id' writes to any memref which escapes (or is an
507// argument to) the block. Returns false otherwise.
509 const Node *node = getNode(id);
510 for (auto *storeOpInst : node->stores) {
511 auto memref = cast<AffineWriteOpInterface>(storeOpInst).getMemRef();
512 auto *op = memref.getDefiningOp();
513 // Return true if 'memref' is a block argument.
514 if (!op)
515 return true;
516 // Return true if any use of 'memref' does not deference it in an affine
517 // way.
518 for (auto *user : memref.getUsers())
519 if (!isa<AffineMapAccessInterface>(*user))
520 return true;
521 }
522 return false;
523}
524
525// Returns true iff there is an edge from node 'srcId' to node 'dstId' which
526// is for 'value' if non-null, or for any value otherwise. Returns false
527// otherwise.
528bool MemRefDependenceGraph::hasEdge(unsigned srcId, unsigned dstId,
529 Value value) const {
530 if (!outEdges.contains(srcId) || !inEdges.contains(dstId)) {
531 return false;
532 }
533 bool hasOutEdge = llvm::any_of(outEdges.lookup(srcId), [=](const Edge &edge) {
534 return edge.id == dstId && (!value || edge.value == value);
535 });
536 bool hasInEdge = llvm::any_of(inEdges.lookup(dstId), [=](const Edge &edge) {
537 return edge.id == srcId && (!value || edge.value == value);
538 });
539 return hasOutEdge && hasInEdge;
540}
541
542// Adds an edge from node 'srcId' to node 'dstId' for 'value'.
543void MemRefDependenceGraph::addEdge(unsigned srcId, unsigned dstId,
544 Value value) {
545 if (!hasEdge(srcId, dstId, value)) {
546 outEdges[srcId].push_back({dstId, value});
547 inEdges[dstId].push_back({srcId, value});
548 if (isa<MemRefType>(value.getType()))
549 memrefEdgeCount[value]++;
550 }
551}
552
553// Removes an edge from node 'srcId' to node 'dstId' for 'value'.
554void MemRefDependenceGraph::removeEdge(unsigned srcId, unsigned dstId,
555 Value value) {
556 assert(inEdges.count(dstId) > 0);
557 assert(outEdges.count(srcId) > 0);
558 if (isa<MemRefType>(value.getType())) {
559 assert(memrefEdgeCount.count(value) > 0);
560 memrefEdgeCount[value]--;
561 }
562 // Remove 'srcId' from 'inEdges[dstId]'.
563 for (auto *it = inEdges[dstId].begin(); it != inEdges[dstId].end(); ++it) {
564 if ((*it).id == srcId && (*it).value == value) {
565 inEdges[dstId].erase(it);
566 break;
567 }
568 }
569 // Remove 'dstId' from 'outEdges[srcId]'.
570 for (auto *it = outEdges[srcId].begin(); it != outEdges[srcId].end(); ++it) {
571 if ((*it).id == dstId && (*it).value == value) {
572 outEdges[srcId].erase(it);
573 break;
574 }
575 }
576}
577
578// Returns true if there is a path in the dependence graph from node 'srcId'
579// to node 'dstId'. Returns false otherwise. `srcId`, `dstId`, and the
580// operations that the edges connected are expected to be from the same block.
582 unsigned dstId) const {
583 // Worklist state is: <node-id, next-output-edge-index-to-visit>
585 worklist.push_back({srcId, 0});
586 Operation *dstOp = getNode(dstId)->op;
587 // Run DFS traversal to see if 'dstId' is reachable from 'srcId'.
588 while (!worklist.empty()) {
589 auto &idAndIndex = worklist.back();
590 // Return true if we have reached 'dstId'.
591 if (idAndIndex.first == dstId)
592 return true;
593 // Pop and continue if node has no out edges, or if all out edges have
594 // already been visited.
595 if (!outEdges.contains(idAndIndex.first) ||
596 idAndIndex.second == outEdges.lookup(idAndIndex.first).size()) {
597 worklist.pop_back();
598 continue;
599 }
600 // Get graph edge to traverse.
601 const Edge edge = outEdges.lookup(idAndIndex.first)[idAndIndex.second];
602 // Increment next output edge index for 'idAndIndex'.
603 ++idAndIndex.second;
604 // Add node at 'edge.id' to the worklist. We don't need to consider
605 // nodes that are "after" dstId in the containing block; one can't have a
606 // path to `dstId` from any of those nodes.
607 bool afterDst = dstOp->isBeforeInBlock(getNode(edge.id)->op);
608 if (!afterDst && edge.id != idAndIndex.first)
609 worklist.push_back({edge.id, 0});
610 }
611 return false;
612}
613
614// Returns the input edge count for node 'id' and 'memref' from src nodes
615// which access 'memref' with a store operation.
617 Value memref) const {
618 unsigned inEdgeCount = 0;
619 for (const Edge &inEdge : inEdges.lookup(id)) {
620 if (inEdge.value == memref) {
621 const Node *srcNode = getNode(inEdge.id);
622 // Only count in edges from 'srcNode' if 'srcNode' accesses 'memref'
623 if (srcNode->getStoreOpCount(memref) > 0)
624 ++inEdgeCount;
625 }
626 }
627 return inEdgeCount;
628}
629
630// Returns the output edge count for node 'id' and 'memref' (if non-null),
631// otherwise returns the total output edge count from node 'id'.
633 Value memref) const {
634 unsigned outEdgeCount = 0;
635 for (const auto &outEdge : outEdges.lookup(id))
636 if (!memref || outEdge.value == memref)
637 ++outEdgeCount;
638 return outEdgeCount;
639}
640
641/// Return all nodes which define SSA values used in node 'id'.
643 unsigned id, DenseSet<unsigned> &definingNodes) const {
644 for (const Edge &edge : inEdges.lookup(id))
645 // By definition of edge, if the edge value is a non-memref value,
646 // then the dependence is between a graph node which defines an SSA value
647 // and another graph node which uses the SSA value.
648 if (!isa<MemRefType>(edge.value.getType()))
649 definingNodes.insert(edge.id);
650}
651
652// Computes and returns an insertion point operation, before which the
653// the fused <srcId, dstId> loop nest can be inserted while preserving
654// dependences. Returns nullptr if no such insertion point is found.
655Operation *
657 unsigned dstId) const {
658 if (!outEdges.contains(srcId))
659 return getNode(dstId)->op;
660
661 // Skip if there is any defining node of 'dstId' that depends on 'srcId'.
662 DenseSet<unsigned> definingNodes;
663 gatherDefiningNodes(dstId, definingNodes);
664 if (llvm::any_of(definingNodes,
665 [&](unsigned id) { return hasDependencePath(srcId, id); })) {
666 LDBG() << "Can't fuse: a defining op with a user in the dst "
667 << "loop has dependence from the src loop";
668 return nullptr;
669 }
670
671 // Build set of insts in range (srcId, dstId) which depend on 'srcId'.
673 for (auto &outEdge : outEdges.lookup(srcId))
674 if (outEdge.id != dstId)
675 srcDepInsts.insert(getNode(outEdge.id)->op);
676
677 // Build set of insts in range (srcId, dstId) on which 'dstId' depends.
679 for (auto &inEdge : inEdges.lookup(dstId))
680 if (inEdge.id != srcId)
681 dstDepInsts.insert(getNode(inEdge.id)->op);
682
683 Operation *srcNodeInst = getNode(srcId)->op;
684 Operation *dstNodeInst = getNode(dstId)->op;
685
686 // Computing insertion point:
687 // *) Walk all operation positions in Block operation list in the
688 // range (src, dst). For each operation 'op' visited in this search:
689 // *) Store in 'firstSrcDepPos' the first position where 'op' has a
690 // dependence edge from 'srcNode'.
691 // *) Store in 'lastDstDepPost' the last position where 'op' has a
692 // dependence edge to 'dstNode'.
693 // *) Compare 'firstSrcDepPos' and 'lastDstDepPost' to determine the
694 // operation insertion point (or return null pointer if no such
695 // insertion point exists: 'firstSrcDepPos' <= 'lastDstDepPos').
697 std::optional<unsigned> firstSrcDepPos;
698 std::optional<unsigned> lastDstDepPos;
699 unsigned pos = 0;
700 for (Block::iterator it = std::next(Block::iterator(srcNodeInst));
701 it != Block::iterator(dstNodeInst); ++it) {
702 Operation *op = &(*it);
703 if (srcDepInsts.count(op) > 0 && firstSrcDepPos == std::nullopt)
704 firstSrcDepPos = pos;
705 if (dstDepInsts.count(op) > 0)
706 lastDstDepPos = pos;
707 depInsts.push_back(op);
708 ++pos;
709 }
710
711 if (firstSrcDepPos.has_value()) {
712 if (lastDstDepPos.has_value()) {
713 if (*firstSrcDepPos <= *lastDstDepPos) {
714 // No valid insertion point exists which preserves dependences.
715 return nullptr;
716 }
717 }
718 // Return the insertion point at 'firstSrcDepPos'.
719 return depInsts[*firstSrcDepPos];
720 }
721 // No dependence targets in range (or only dst deps in range), return
722 // 'dstNodInst' insertion point.
723 return dstNodeInst;
724}
725
726// Updates edge mappings from node 'srcId' to node 'dstId' after fusing them,
727// taking into account that:
728// *) if 'removeSrcId' is true, 'srcId' will be removed after fusion,
729// *) memrefs in 'privateMemRefs' has been replaced in node at 'dstId' by a
730// private memref.
731void MemRefDependenceGraph::updateEdges(unsigned srcId, unsigned dstId,
732 const DenseSet<Value> &privateMemRefs,
733 bool removeSrcId) {
734 // For each edge in 'inEdges[srcId]': add new edge remapping to 'dstId'.
735 if (inEdges.count(srcId) > 0) {
736 SmallVector<Edge, 2> oldInEdges = inEdges[srcId];
737 for (auto &inEdge : oldInEdges) {
738 // Add edge from 'inEdge.id' to 'dstId' if it's not a private memref.
739 if (!privateMemRefs.contains(inEdge.value))
740 addEdge(inEdge.id, dstId, inEdge.value);
741 }
742 }
743 // For each edge in 'outEdges[srcId]': remove edge from 'srcId' to 'dstId'.
744 // If 'srcId' is going to be removed, remap all the out edges to 'dstId'.
745 if (outEdges.count(srcId) > 0) {
746 SmallVector<Edge, 2> oldOutEdges = outEdges[srcId];
747 for (auto &outEdge : oldOutEdges) {
748 // Remove any out edges from 'srcId' to 'dstId' across memrefs.
749 if (outEdge.id == dstId)
750 removeEdge(srcId, outEdge.id, outEdge.value);
751 else if (removeSrcId) {
752 addEdge(dstId, outEdge.id, outEdge.value);
753 removeEdge(srcId, outEdge.id, outEdge.value);
754 }
755 }
756 }
757 // Remove any edges in 'inEdges[dstId]' on 'oldMemRef' (which is being
758 // replaced by a private memref). These edges could come from nodes
759 // other than 'srcId' which were removed in the previous step.
760 if (inEdges.count(dstId) > 0 && !privateMemRefs.empty()) {
761 SmallVector<Edge, 2> oldInEdges = inEdges[dstId];
762 for (auto &inEdge : oldInEdges)
763 if (privateMemRefs.count(inEdge.value) > 0)
764 removeEdge(inEdge.id, dstId, inEdge.value);
765 }
766}
767
768// Update edge mappings for nodes 'sibId' and 'dstId' to reflect fusion
769// of sibling node 'sibId' into node 'dstId'.
770void MemRefDependenceGraph::updateEdges(unsigned sibId, unsigned dstId) {
771 // For each edge in 'inEdges[sibId]':
772 // *) Add new edge from source node 'inEdge.id' to 'dstNode'.
773 // *) Remove edge from source node 'inEdge.id' to 'sibNode'.
774 if (inEdges.count(sibId) > 0) {
775 SmallVector<Edge, 2> oldInEdges = inEdges[sibId];
776 for (auto &inEdge : oldInEdges) {
777 addEdge(inEdge.id, dstId, inEdge.value);
778 removeEdge(inEdge.id, sibId, inEdge.value);
779 }
780 }
781
782 // For each edge in 'outEdges[sibId]' to node 'id'
783 // *) Add new edge from 'dstId' to 'outEdge.id'.
784 // *) Remove edge from 'sibId' to 'outEdge.id'.
785 if (outEdges.count(sibId) > 0) {
786 SmallVector<Edge, 2> oldOutEdges = outEdges[sibId];
787 for (auto &outEdge : oldOutEdges) {
788 addEdge(dstId, outEdge.id, outEdge.value);
789 removeEdge(sibId, outEdge.id, outEdge.value);
790 }
791 }
792}
793
794// Adds ops in 'loads' and 'stores' to node at 'id'.
797 ArrayRef<Operation *> memrefLoads,
798 ArrayRef<Operation *> memrefStores,
799 ArrayRef<Operation *> memrefFrees) {
800 Node *node = getNode(id);
801 llvm::append_range(node->loads, loads);
802 llvm::append_range(node->stores, stores);
803 llvm::append_range(node->memrefLoads, memrefLoads);
804 llvm::append_range(node->memrefStores, memrefStores);
805 llvm::append_range(node->memrefFrees, memrefFrees);
806}
807
809 Node *node = getNode(id);
810 node->loads.clear();
811 node->stores.clear();
812}
813
814// Calls 'callback' for each input edge incident to node 'id' which carries a
815// memref dependence.
817 unsigned id, const std::function<void(Edge)> &callback) {
818 if (inEdges.count(id) > 0)
819 forEachMemRefEdge(inEdges.at(id), callback);
820}
821
822// Calls 'callback' for each output edge from node 'id' which carries a
823// memref dependence.
825 unsigned id, const std::function<void(Edge)> &callback) {
826 if (outEdges.count(id) > 0)
827 forEachMemRefEdge(outEdges.at(id), callback);
828}
829
830// Calls 'callback' for each edge in 'edges' which carries a memref
831// dependence.
833 ArrayRef<Edge> edges, const std::function<void(Edge)> &callback) {
834 for (const auto &edge : edges) {
835 // Skip if 'edge' is not a memref dependence edge.
836 if (!isa<MemRefType>(edge.value.getType()))
837 continue;
838 assert(nodes.count(edge.id) > 0);
839 // Visit current input edge 'edge'.
840 callback(edge);
841 }
842}
843
845 os << "\nMemRefDependenceGraph\n";
846 os << "\nNodes:\n";
847 for (const auto &idAndNode : nodes) {
848 os << "Node: " << idAndNode.first << "\n";
849 auto it = inEdges.find(idAndNode.first);
850 if (it != inEdges.end()) {
851 for (const auto &e : it->second)
852 os << " InEdge: " << e.id << " " << e.value << "\n";
853 }
854 it = outEdges.find(idAndNode.first);
855 if (it != outEdges.end()) {
856 for (const auto &e : it->second)
857 os << " OutEdge: " << e.id << " " << e.value << "\n";
858 }
859 }
860}
861
864 auto *currOp = op.getParentOp();
865 AffineForOp currAffineForOp;
866 // Traverse up the hierarchy collecting all 'affine.for' operation while
867 // skipping over 'affine.if' operations.
868 while (currOp && !currOp->hasTrait<OpTrait::AffineScope>()) {
869 if (auto currAffineForOp = dyn_cast<AffineForOp>(currOp))
870 loops->push_back(currAffineForOp);
871 currOp = currOp->getParentOp();
872 }
873 std::reverse(loops->begin(), loops->end());
874}
875
878 ops->clear();
879 Operation *currOp = op.getParentOp();
880
881 // Traverse up the hierarchy collecting all `affine.for`, `affine.if`, and
882 // affine.parallel operations.
883 while (currOp && !currOp->hasTrait<OpTrait::AffineScope>()) {
884 if (isa<AffineIfOp, AffineForOp, AffineParallelOp>(currOp))
885 ops->push_back(currOp);
886 currOp = currOp->getParentOp();
887 }
888 std::reverse(ops->begin(), ops->end());
889}
890
891// Populates 'cst' with FlatAffineValueConstraints which represent original
892// domain of the loop bounds that define 'ivs'.
894 FlatAffineValueConstraints &cst) const {
895 assert(!ivs.empty() && "Cannot have a slice without its IVs");
896 cst = FlatAffineValueConstraints(/*numDims=*/ivs.size(), /*numSymbols=*/0,
897 /*numLocals=*/0, ivs);
898 for (Value iv : ivs) {
899 AffineForOp loop = getForInductionVarOwner(iv);
900 assert(loop && "Expected affine for");
901 if (failed(cst.addAffineForOpDomain(loop)))
902 return failure();
903 }
904 return success();
905}
906
907// Populates 'cst' with FlatAffineValueConstraints which represent slice bounds.
908LogicalResult
910 assert(!lbOperands.empty());
911 // Adds src 'ivs' as dimension variables in 'cst'.
912 unsigned numDims = ivs.size();
913 // Adds operands (dst ivs and symbols) as symbols in 'cst'.
914 unsigned numSymbols = lbOperands[0].size();
915
917 // Append 'ivs' then 'operands' to 'values'.
918 values.append(lbOperands[0].begin(), lbOperands[0].end());
919 *cst = FlatAffineValueConstraints(numDims, numSymbols, 0, values);
920
921 // Add loop bound constraints for values which are loop IVs of the destination
922 // of fusion and equality constraints for symbols which are constants.
923 for (unsigned i = numDims, end = values.size(); i < end; ++i) {
924 Value value = values[i];
925 assert(cst->containsVar(value) && "value expected to be present");
926 if (isValidSymbol(value)) {
927 // Check if the symbol is a constant.
928 if (std::optional<int64_t> cOp = getConstantIntValue(value))
929 cst->addBound(BoundType::EQ, value, cOp.value());
930 } else if (auto loop = getForInductionVarOwner(value)) {
931 if (failed(cst->addAffineForOpDomain(loop)))
932 return failure();
933 }
934 }
935
936 // Add slices bounds on 'ivs' using maps 'lbs'/'ubs' with 'lbOperands[0]'
937 LogicalResult ret = cst->addSliceBounds(ivs, lbs, ubs, lbOperands[0]);
938 assert(succeeded(ret) &&
939 "should not fail as we never have semi-affine slice maps");
940 (void)ret;
941 return success();
942}
943
944// Clears state bounds and operand state.
946 lbs.clear();
947 ubs.clear();
948 lbOperands.clear();
949 ubOperands.clear();
950}
951
953 llvm::errs() << "\tIVs:\n";
954 for (Value iv : ivs)
955 llvm::errs() << "\t\t" << iv << "\n";
956
957 llvm::errs() << "\tLBs:\n";
958 for (auto en : llvm::enumerate(lbs)) {
959 llvm::errs() << "\t\t" << en.value() << "\n";
960 llvm::errs() << "\t\tOperands:\n";
961 for (Value lbOp : lbOperands[en.index()])
962 llvm::errs() << "\t\t\t" << lbOp << "\n";
963 }
964
965 llvm::errs() << "\tUBs:\n";
966 for (auto en : llvm::enumerate(ubs)) {
967 llvm::errs() << "\t\t" << en.value() << "\n";
968 llvm::errs() << "\t\tOperands:\n";
969 for (Value ubOp : ubOperands[en.index()])
970 llvm::errs() << "\t\t\t" << ubOp << "\n";
971 }
972}
973
974/// Fast check to determine if the computation slice is maximal. Returns true if
975/// each slice dimension maps to an existing dst dimension and both the src
976/// and the dst loops for those dimensions have the same bounds. Returns false
977/// if both the src and the dst loops don't have the same bounds. Returns
978/// std::nullopt if none of the above can be proven.
979std::optional<bool> ComputationSliceState::isSliceMaximalFastCheck() const {
980 assert(lbs.size() == ubs.size() && !lbs.empty() && !ivs.empty() &&
981 "Unexpected number of lbs, ubs and ivs in slice");
982
983 for (unsigned i = 0, end = lbs.size(); i < end; ++i) {
984 AffineMap lbMap = lbs[i];
985 AffineMap ubMap = ubs[i];
986
987 // Check if this slice is just an equality along this dimension.
988 if (!lbMap || !ubMap || lbMap.getNumResults() != 1 ||
989 ubMap.getNumResults() != 1 ||
990 lbMap.getResult(0) + 1 != ubMap.getResult(0) ||
991 // The condition above will be true for maps describing a single
992 // iteration (e.g., lbMap.getResult(0) = 0, ubMap.getResult(0) = 1).
993 // Make sure we skip those cases by checking that the lb result is not
994 // just a constant.
995 isa<AffineConstantExpr>(lbMap.getResult(0)))
996 return std::nullopt;
997
998 // Limited support: we expect the lb result to be just a loop dimension for
999 // now.
1000 AffineDimExpr result = dyn_cast<AffineDimExpr>(lbMap.getResult(0));
1001 if (!result)
1002 return std::nullopt;
1003
1004 // Retrieve dst loop bounds.
1005 AffineForOp dstLoop =
1006 getForInductionVarOwner(lbOperands[i][result.getPosition()]);
1007 if (!dstLoop)
1008 return std::nullopt;
1009 AffineMap dstLbMap = dstLoop.getLowerBoundMap();
1010 AffineMap dstUbMap = dstLoop.getUpperBoundMap();
1011
1012 // Retrieve src loop bounds.
1013 AffineForOp srcLoop = getForInductionVarOwner(ivs[i]);
1014 assert(srcLoop && "Expected affine for");
1015 AffineMap srcLbMap = srcLoop.getLowerBoundMap();
1016 AffineMap srcUbMap = srcLoop.getUpperBoundMap();
1017
1018 // Limited support: we expect simple src and dst loops with a single
1019 // constant component per bound for now.
1020 if (srcLbMap.getNumResults() != 1 || srcUbMap.getNumResults() != 1 ||
1021 dstLbMap.getNumResults() != 1 || dstUbMap.getNumResults() != 1)
1022 return std::nullopt;
1023
1024 AffineExpr srcLbResult = srcLbMap.getResult(0);
1025 AffineExpr dstLbResult = dstLbMap.getResult(0);
1026 AffineExpr srcUbResult = srcUbMap.getResult(0);
1027 AffineExpr dstUbResult = dstUbMap.getResult(0);
1028 if (!isa<AffineConstantExpr>(srcLbResult) ||
1029 !isa<AffineConstantExpr>(srcUbResult) ||
1030 !isa<AffineConstantExpr>(dstLbResult) ||
1031 !isa<AffineConstantExpr>(dstUbResult))
1032 return std::nullopt;
1033
1034 // Check if src and dst loop bounds are the same. If not, we can guarantee
1035 // that the slice is not maximal.
1036 if (srcLbResult != dstLbResult || srcUbResult != dstUbResult ||
1037 srcLoop.getStep() != dstLoop.getStep())
1038 return false;
1039 }
1040
1041 return true;
1042}
1043
1044/// Returns true if it is deterministically verified that the original iteration
1045/// space of the slice is contained within the new iteration space that is
1046/// created after fusing 'this' slice into its destination.
1047std::optional<bool> ComputationSliceState::isSliceValid() const {
1048 // Fast check to determine if the slice is valid. If the following conditions
1049 // are verified to be true, slice is declared valid by the fast check:
1050 // 1. Each slice loop is a single iteration loop bound in terms of a single
1051 // destination loop IV.
1052 // 2. Loop bounds of the destination loop IV (from above) and those of the
1053 // source loop IV are exactly the same.
1054 // If the fast check is inconclusive or false, we proceed with a more
1055 // expensive analysis.
1056 // TODO: Store the result of the fast check, as it might be used again in
1057 // `canRemoveSrcNodeAfterFusion`.
1058 std::optional<bool> isValidFastCheck = isSliceMaximalFastCheck();
1059 if (isValidFastCheck && *isValidFastCheck)
1060 return true;
1061
1062 // Create constraints for the source loop nest using which slice is computed.
1063 FlatAffineValueConstraints srcConstraints;
1064 // TODO: Store the source's domain to avoid computation at each depth.
1065 if (failed(getSourceAsConstraints(srcConstraints))) {
1066 LDBG() << "Unable to compute source's domain";
1067 return std::nullopt;
1068 }
1069 // TODO: Handle local vars in the source domains while using the 'projectOut'
1070 // utility below. Currently, aligning is not done assuming that there will be
1071 // no local vars in the source domain.
1072 if (srcConstraints.getNumLocalVars() != 0) {
1073 LDBG() << "Cannot handle locals in source domain";
1074 return std::nullopt;
1075 }
1076
1077 // Create constraints for the slice loop nest that would be created if the
1078 // fusion succeeds.
1079 FlatAffineValueConstraints sliceConstraints;
1080 if (failed(getAsConstraints(&sliceConstraints))) {
1081 LDBG() << "Unable to compute slice's domain";
1082 return std::nullopt;
1083 }
1084
1085 // Projecting out every dimension other than the 'ivs' to express slice's
1086 // domain completely in terms of source's IVs.
1087 sliceConstraints.projectOut(ivs.size(),
1088 sliceConstraints.getNumVars() - ivs.size());
1089 srcConstraints.projectOut(ivs.size(),
1090 srcConstraints.getNumVars() - ivs.size());
1091
1092 LDBG() << "Domain of the source of the slice:\n"
1093 << "Source constraints:" << srcConstraints
1094 << "\nDomain of the slice if this fusion succeeds "
1095 << "(expressed in terms of its source's IVs):\n"
1096 << "Slice constraints:" << sliceConstraints;
1097
1098 // TODO: Store 'srcSet' to avoid recalculating for each depth.
1099 PresburgerSet srcSet(srcConstraints);
1100 PresburgerSet sliceSet(sliceConstraints);
1101 PresburgerSet diffSet = sliceSet.subtract(srcSet);
1102
1103 if (!diffSet.isIntegerEmpty()) {
1104 LDBG() << "Incorrect slice";
1105 return false;
1106 }
1107 return true;
1108}
1109
1110/// Returns true if the computation slice encloses all the iterations of the
1111/// sliced loop nest. Returns false if it does not. Returns std::nullopt if it
1112/// cannot determine if the slice is maximal or not.
1113std::optional<bool> ComputationSliceState::isMaximal() const {
1114 // Fast check to determine if the computation slice is maximal. If the result
1115 // is inconclusive, we proceed with a more expensive analysis.
1116 std::optional<bool> isMaximalFastCheck = isSliceMaximalFastCheck();
1117 if (isMaximalFastCheck)
1118 return isMaximalFastCheck;
1119
1120 // Create constraints for the src loop nest being sliced.
1121 FlatAffineValueConstraints srcConstraints(/*numDims=*/ivs.size(),
1122 /*numSymbols=*/0,
1123 /*numLocals=*/0, ivs);
1124 for (Value iv : ivs) {
1125 AffineForOp loop = getForInductionVarOwner(iv);
1126 assert(loop && "Expected affine for");
1127 if (failed(srcConstraints.addAffineForOpDomain(loop)))
1128 return std::nullopt;
1129 }
1130
1131 // Create constraints for the slice using the dst loop nest information. We
1132 // retrieve existing dst loops from the lbOperands.
1133 SmallVector<Value> consumerIVs;
1134 for (Value lbOp : lbOperands[0])
1135 if (getForInductionVarOwner(lbOp))
1136 consumerIVs.push_back(lbOp);
1137
1138 // Add empty IV Values for those new loops that are not equalities and,
1139 // therefore, are not yet materialized in the IR.
1140 for (int i = consumerIVs.size(), end = ivs.size(); i < end; ++i)
1141 consumerIVs.push_back(Value());
1142
1143 FlatAffineValueConstraints sliceConstraints(/*numDims=*/consumerIVs.size(),
1144 /*numSymbols=*/0,
1145 /*numLocals=*/0, consumerIVs);
1146
1147 if (failed(sliceConstraints.addDomainFromSliceMaps(lbs, ubs, lbOperands[0])))
1148 return std::nullopt;
1149
1150 if (srcConstraints.getNumDimVars() != sliceConstraints.getNumDimVars())
1151 // Constraint dims are different. The integer set difference can't be
1152 // computed so we don't know if the slice is maximal.
1153 return std::nullopt;
1154
1155 // Compute the difference between the src loop nest and the slice integer
1156 // sets.
1157 PresburgerSet srcSet(srcConstraints);
1158 PresburgerSet sliceSet(sliceConstraints);
1159 PresburgerSet diffSet = srcSet.subtract(sliceSet);
1160 return diffSet.isIntegerEmpty();
1161}
1162
1163unsigned MemRefRegion::getRank() const {
1164 return cast<MemRefType>(memref.getType()).getRank();
1165}
1166
1169 auto memRefType = cast<MemRefType>(memref.getType());
1170 MLIRContext *context = memref.getContext();
1171 unsigned rank = memRefType.getRank();
1172 if (shape)
1173 shape->reserve(rank);
1174
1175 assert(rank == cst.getNumDimVars() && "inconsistent memref region");
1176
1177 // Use a copy of the region constraints that has upper/lower bounds for each
1178 // memref dimension with static size added to guard against potential
1179 // over-approximation from projection or union bounding box. We may not add
1180 // this on the region itself since they might just be redundant constraints
1181 // that will need non-trivials means to eliminate.
1182 FlatLinearValueConstraints cstWithShapeBounds(cst);
1183 for (unsigned r = 0; r < rank; r++) {
1184 cstWithShapeBounds.addBound(BoundType::LB, r, 0);
1185 int64_t dimSize = memRefType.getDimSize(r);
1186 if (ShapedType::isDynamic(dimSize))
1187 continue;
1188 cstWithShapeBounds.addBound(BoundType::UB, r, dimSize - 1);
1189 }
1190
1191 // Find a constant upper bound on the extent of this memref region along
1192 // each dimension.
1193 int64_t numElements = 1;
1194 int64_t diffConstant;
1195 for (unsigned d = 0; d < rank; d++) {
1196 AffineMap lb;
1197 std::optional<int64_t> diff =
1198 cstWithShapeBounds.getConstantBoundOnDimSize(context, d, &lb);
1199 if (diff.has_value()) {
1200 diffConstant = *diff;
1201 assert(diffConstant >= 0 && "dim size bound cannot be negative");
1202 } else {
1203 // If no constant bound is found, then it can always be bound by the
1204 // memref's dim size if the latter has a constant size along this dim.
1205 auto dimSize = memRefType.getDimSize(d);
1206 if (ShapedType::isDynamic(dimSize))
1207 return std::nullopt;
1208 diffConstant = dimSize;
1209 // Lower bound becomes 0.
1210 lb = AffineMap::get(/*dimCount=*/0, cstWithShapeBounds.getNumSymbolVars(),
1211 /*result=*/getAffineConstantExpr(0, context));
1212 }
1213 numElements *= diffConstant;
1214 // Populate outputs if available.
1215 if (lbs)
1216 lbs->push_back(lb);
1217 if (shape)
1218 shape->push_back(diffConstant);
1219 }
1220 return numElements;
1221}
1222
1224 AffineMap &ubMap) const {
1225 assert(pos < cst.getNumDimVars() && "invalid position");
1226 auto memRefType = cast<MemRefType>(memref.getType());
1227 unsigned rank = memRefType.getRank();
1228
1229 assert(rank == cst.getNumDimVars() && "inconsistent memref region");
1230
1231 auto boundPairs = cst.getLowerAndUpperBound(
1232 pos, /*offset=*/0, /*num=*/rank, cst.getNumDimAndSymbolVars(),
1233 /*localExprs=*/{}, memRefType.getContext());
1234 lbMap = boundPairs.first;
1235 ubMap = boundPairs.second;
1236 assert(lbMap && "lower bound for a region must exist");
1237 assert(ubMap && "upper bound for a region must exist");
1238 assert(lbMap.getNumInputs() == cst.getNumDimAndSymbolVars() - rank);
1239 assert(ubMap.getNumInputs() == cst.getNumDimAndSymbolVars() - rank);
1240}
1241
1243 assert(memref == other.memref);
1244 return cst.unionBoundingBox(*other.getConstraints());
1245}
1246
1247/// Computes the memory region accessed by this memref with the region
1248/// represented as constraints symbolic/parametric in 'loopDepth' loops
1249/// surrounding opInst and any additional Function symbols.
1250// For example, the memref region for this load operation at loopDepth = 1 will
1251// be as below:
1252//
1253// affine.for %i = 0 to 32 {
1254// affine.for %ii = %i to (d0) -> (d0 + 8) (%i) {
1255// load %A[%ii]
1256// }
1257// }
1258//
1259// region: {memref = %A, write = false, {%i <= m0 <= %i + 7} }
1260// The last field is a 2-d FlatAffineValueConstraints symbolic in %i.
1261//
1262// TODO: extend this to any other memref dereferencing ops
1263// (dma_start, dma_wait).
1264LogicalResult MemRefRegion::compute(Operation *op, unsigned loopDepth,
1265 const ComputationSliceState *sliceState,
1266 bool addMemRefDimBounds, bool dropLocalVars,
1267 bool dropOuterIvs) {
1268 assert((isa<AffineReadOpInterface, AffineWriteOpInterface>(op)) &&
1269 "affine read/write op expected");
1270
1271 MemRefAccess access(op);
1272 memref = access.memref;
1273 write = access.isStore();
1274
1275 unsigned rank = access.getRank();
1276
1277 LDBG() << "MemRefRegion::compute: " << *op << " depth: " << loopDepth;
1278
1279 // 0-d memrefs.
1280 if (rank == 0) {
1282 getAffineIVs(*op, ivs);
1283 assert(loopDepth <= ivs.size() && "invalid 'loopDepth'");
1284 // The first 'loopDepth' IVs are symbols for this region.
1285 ivs.resize(loopDepth);
1286 // A 0-d memref has a 0-d region.
1287 cst = FlatAffineValueConstraints(rank, loopDepth, /*numLocals=*/0, ivs);
1288 return success();
1289 }
1290
1291 // Build the constraints for this region.
1292 AffineValueMap accessValueMap;
1293 access.getAccessMap(&accessValueMap);
1294 AffineMap accessMap = accessValueMap.getAffineMap();
1295
1296 unsigned numDims = accessMap.getNumDims();
1297 unsigned numSymbols = accessMap.getNumSymbols();
1298 unsigned numOperands = accessValueMap.getNumOperands();
1299 // Merge operands with slice operands.
1300 SmallVector<Value, 4> operands;
1301 operands.resize(numOperands);
1302 for (unsigned i = 0; i < numOperands; ++i)
1303 operands[i] = accessValueMap.getOperand(i);
1304
1305 if (sliceState != nullptr) {
1306 operands.reserve(operands.size() + sliceState->lbOperands[0].size());
1307 // Append slice operands to 'operands' as symbols.
1308 for (auto extraOperand : sliceState->lbOperands[0]) {
1309 if (!llvm::is_contained(operands, extraOperand)) {
1310 operands.push_back(extraOperand);
1311 numSymbols++;
1312 }
1313 }
1314 }
1315 // We'll first associate the dims and symbols of the access map to the dims
1316 // and symbols resp. of cst. This will change below once cst is
1317 // fully constructed out.
1318 cst = FlatAffineValueConstraints(numDims, numSymbols, 0, operands);
1319
1320 // Add equality constraints.
1321 // Add inequalities for loop lower/upper bounds.
1322 for (unsigned i = 0; i < numDims + numSymbols; ++i) {
1323 auto operand = operands[i];
1324 if (auto affineFor = getForInductionVarOwner(operand)) {
1325 // Note that cst can now have more dimensions than accessMap if the
1326 // bounds expressions involve outer loops or other symbols.
1327 // TODO: rewrite this to use getInstIndexSet; this way
1328 // conditionals will be handled when the latter supports it.
1329 if (failed(cst.addAffineForOpDomain(affineFor)))
1330 return failure();
1331 } else if (auto parallelOp = getAffineParallelInductionVarOwner(operand)) {
1332 if (failed(cst.addAffineParallelOpDomain(parallelOp)))
1333 return failure();
1334 } else if (isValidSymbol(operand)) {
1335 // Check if the symbol is a constant.
1336 Value symbol = operand;
1337 if (auto constVal = getConstantIntValue(symbol))
1338 cst.addBound(BoundType::EQ, symbol, constVal.value());
1339 } else {
1340 LDBG() << "unknown affine dimensional value";
1341 return failure();
1342 }
1343 }
1344
1345 // Add lower/upper bounds on loop IVs using bounds from 'sliceState'.
1346 if (sliceState != nullptr) {
1347 // Add dim and symbol slice operands.
1348 for (auto operand : sliceState->lbOperands[0]) {
1349 if (failed(cst.addInductionVarOrTerminalSymbol(operand)))
1350 return failure();
1351 }
1352 // Add upper/lower bounds from 'sliceState' to 'cst'.
1353 LogicalResult ret =
1354 cst.addSliceBounds(sliceState->ivs, sliceState->lbs, sliceState->ubs,
1355 sliceState->lbOperands[0]);
1356 assert(succeeded(ret) &&
1357 "should not fail as we never have semi-affine slice maps");
1358 (void)ret;
1359 }
1360
1361 // Add access function equalities to connect loop IVs to data dimensions.
1362 if (failed(cst.composeMap(&accessValueMap))) {
1363 op->emitError("getMemRefRegion: compose affine map failed");
1364 LDBG() << "Access map: " << accessValueMap.getAffineMap();
1365 return failure();
1366 }
1367
1368 // Set all variables appearing after the first 'rank' variables as
1369 // symbolic variables - so that the ones corresponding to the memref
1370 // dimensions are the dimensional variables for the memref region.
1371 cst.setDimSymbolSeparation(cst.getNumDimAndSymbolVars() - rank);
1372
1373 // Eliminate any loop IVs other than the outermost 'loopDepth' IVs, on which
1374 // this memref region is symbolic.
1375 SmallVector<Value, 4> enclosingIVs;
1376 getAffineIVs(*op, enclosingIVs);
1377 assert(loopDepth <= enclosingIVs.size() && "invalid loop depth");
1378 enclosingIVs.resize(loopDepth);
1380 cst.getValues(cst.getNumDimVars(), cst.getNumDimAndSymbolVars(), &vars);
1381 for (auto en : llvm::enumerate(vars)) {
1382 if ((isAffineInductionVar(en.value())) &&
1383 !llvm::is_contained(enclosingIVs, en.value())) {
1384 if (dropOuterIvs) {
1385 cst.projectOut(en.value());
1386 } else {
1387 unsigned varPosition;
1388 cst.findVar(en.value(), &varPosition);
1389 auto varKind = cst.getVarKindAt(varPosition);
1390 varPosition -= cst.getNumDimVars();
1391 cst.convertToLocal(varKind, varPosition, varPosition + 1);
1392 }
1393 }
1394 }
1395
1396 // Project out any local variables (these would have been added for any
1397 // mod/divs) if specified.
1398 if (dropLocalVars)
1399 cst.projectOut(cst.getNumDimAndSymbolVars(), cst.getNumLocalVars());
1400
1401 // Constant fold any symbolic variables.
1402 cst.constantFoldVarRange(/*pos=*/cst.getNumDimVars(),
1403 /*num=*/cst.getNumSymbolVars());
1404
1405 assert(cst.getNumDimVars() == rank && "unexpected MemRefRegion format");
1406
1407 // Add upper/lower bounds for each memref dimension with static size
1408 // to guard against potential over-approximation from projection.
1409 // TODO: Support dynamic memref dimensions.
1410 if (addMemRefDimBounds) {
1411 auto memRefType = cast<MemRefType>(memref.getType());
1412 for (unsigned r = 0; r < rank; r++) {
1413 cst.addBound(BoundType::LB, /*pos=*/r, /*value=*/0);
1414 if (memRefType.isDynamicDim(r))
1415 continue;
1416 cst.addBound(BoundType::UB, /*pos=*/r, memRefType.getDimSize(r) - 1);
1417 }
1418 }
1419 cst.removeTrivialRedundancy();
1420
1421 LDBG() << "Memory region: " << cst;
1422 return success();
1423}
1424
1425std::optional<int64_t>
1427 auto elementType = memRefType.getElementType();
1428
1429 unsigned sizeInBits;
1430 if (elementType.isIntOrFloat()) {
1431 sizeInBits = elementType.getIntOrFloatBitWidth();
1432 } else if (auto vectorType = dyn_cast<VectorType>(elementType)) {
1433 if (vectorType.getElementType().isIntOrFloat())
1434 sizeInBits =
1435 vectorType.getElementTypeBitWidth() * vectorType.getNumElements();
1436 else
1437 return std::nullopt;
1438 } else {
1439 return std::nullopt;
1440 }
1441 return llvm::divideCeil(sizeInBits, 8);
1442}
1443
1444// Returns the size of the region.
1445std::optional<int64_t> MemRefRegion::getRegionSize() {
1446 auto memRefType = cast<MemRefType>(memref.getType());
1447
1448 if (!memRefType.getLayout().isIdentity()) {
1449 LDBG() << "Non-identity layout map not yet supported";
1450 return false;
1451 }
1452
1453 // Compute the extents of the buffer.
1454 std::optional<int64_t> numElements = getConstantBoundingSizeAndShape();
1455 if (!numElements) {
1456 LDBG() << "Dynamic shapes not yet supported";
1457 return std::nullopt;
1458 }
1459 auto eltSize = getMemRefIntOrFloatEltSizeInBytes(memRefType);
1460 if (!eltSize)
1461 return std::nullopt;
1462 return *eltSize * *numElements;
1463}
1464
1465/// Returns the size of memref data in bytes if it's statically shaped,
1466/// std::nullopt otherwise. If the element of the memref has vector type, takes
1467/// into account size of the vector as well.
1468// TODO: improve/complete this when we have target data.
1469std::optional<uint64_t>
1471 if (!memRefType.hasStaticShape())
1472 return std::nullopt;
1473 auto elementType = memRefType.getElementType();
1474 if (!elementType.isIntOrFloat() && !isa<VectorType>(elementType))
1475 return std::nullopt;
1476
1477 auto sizeInBytes = getMemRefIntOrFloatEltSizeInBytes(memRefType);
1478 if (!sizeInBytes)
1479 return std::nullopt;
1480 for (unsigned i = 0, e = memRefType.getRank(); i < e; i++) {
1481 sizeInBytes = *sizeInBytes * memRefType.getDimSize(i);
1482 }
1483 return sizeInBytes;
1484}
1485
1486template <typename LoadOrStoreOp>
1487LogicalResult mlir::affine::boundCheckLoadOrStoreOp(LoadOrStoreOp loadOrStoreOp,
1488 bool emitError) {
1489 static_assert(llvm::is_one_of<LoadOrStoreOp, AffineReadOpInterface,
1490 AffineWriteOpInterface>::value,
1491 "argument should be either a AffineReadOpInterface or a "
1492 "AffineWriteOpInterface");
1493
1494 Operation *op = loadOrStoreOp.getOperation();
1495 MemRefRegion region(op->getLoc());
1496 if (failed(region.compute(op, /*loopDepth=*/0, /*sliceState=*/nullptr,
1497 /*addMemRefDimBounds=*/false)))
1498 return success();
1499
1500 LDBG() << "Memory region: " << region.getConstraints();
1501
1502 bool outOfBounds = false;
1503 unsigned rank = loadOrStoreOp.getMemRefType().getRank();
1504
1505 // For each dimension, check for out of bounds.
1506 for (unsigned r = 0; r < rank; r++) {
1507 FlatAffineValueConstraints ucst(*region.getConstraints());
1508
1509 // Intersect memory region with constraint capturing out of bounds (both out
1510 // of upper and out of lower), and check if the constraint system is
1511 // feasible. If it is, there is at least one point out of bounds.
1512 SmallVector<int64_t, 4> ineq(rank + 1, 0);
1513 int64_t dimSize = loadOrStoreOp.getMemRefType().getDimSize(r);
1514 // TODO: handle dynamic dim sizes.
1515 if (dimSize == -1)
1516 continue;
1517
1518 // Check for overflow: d_i >= memref dim size.
1519 ucst.addBound(BoundType::LB, r, dimSize);
1520 outOfBounds = !ucst.isEmpty();
1521 if (outOfBounds && emitError) {
1522 loadOrStoreOp.emitOpError()
1523 << "memref out of upper bound access along dimension #" << (r + 1);
1524 }
1525
1526 // Check for a negative index.
1527 FlatAffineValueConstraints lcst(*region.getConstraints());
1528 llvm::fill(ineq, 0);
1529 // d_i <= -1;
1530 lcst.addBound(BoundType::UB, r, -1);
1531 outOfBounds = !lcst.isEmpty();
1532 if (outOfBounds && emitError) {
1533 loadOrStoreOp.emitOpError()
1534 << "memref out of lower bound access along dimension #" << (r + 1);
1535 }
1536 }
1537 return failure(outOfBounds);
1538}
1539
1540// Explicitly instantiate the template so that the compiler knows we need them!
1541template LogicalResult
1542mlir::affine::boundCheckLoadOrStoreOp(AffineReadOpInterface loadOp,
1543 bool emitError);
1544template LogicalResult
1545mlir::affine::boundCheckLoadOrStoreOp(AffineWriteOpInterface storeOp,
1546 bool emitError);
1547
1548// Returns in 'positions' the Block positions of 'op' in each ancestor
1549// Block from the Block containing operation, stopping at 'limitBlock'.
1550static void findInstPosition(Operation *op, Block *limitBlock,
1551 SmallVectorImpl<unsigned> *positions) {
1552 Block *block = op->getBlock();
1553 while (block != limitBlock) {
1554 // FIXME: This algorithm is unnecessarily O(n) and should be improved to not
1555 // rely on linear scans.
1556 int instPosInBlock = std::distance(block->begin(), op->getIterator());
1557 positions->push_back(instPosInBlock);
1558 op = block->getParentOp();
1559 block = op->getBlock();
1560 }
1561 std::reverse(positions->begin(), positions->end());
1562}
1563
1564// Returns the Operation in a possibly nested set of Blocks, where the
1565// position of the operation is represented by 'positions', which has a
1566// Block position for each level of nesting.
1568 unsigned level, Block *block) {
1569 unsigned i = 0;
1570 for (auto &op : *block) {
1571 if (i != positions[level]) {
1572 ++i;
1573 continue;
1574 }
1575 if (level == positions.size() - 1)
1576 return &op;
1577 if (auto childAffineForOp = dyn_cast<AffineForOp>(op))
1578 return getInstAtPosition(positions, level + 1,
1579 childAffineForOp.getBody());
1580
1581 for (auto &region : op.getRegions()) {
1582 for (auto &b : region)
1583 if (auto *ret = getInstAtPosition(positions, level + 1, &b))
1584 return ret;
1585 }
1586 return nullptr;
1587 }
1588 return nullptr;
1589}
1590
1591// Adds loop IV bounds to 'cst' for loop IVs not found in 'ivs'.
1594 for (unsigned i = 0, e = cst->getNumDimVars(); i < e; ++i) {
1595 auto value = cst->getValue(i);
1596 if (ivs.count(value) == 0) {
1597 assert(isAffineForInductionVar(value));
1598 auto loop = getForInductionVarOwner(value);
1599 if (failed(cst->addAffineForOpDomain(loop)))
1600 return failure();
1601 }
1602 }
1603 return success();
1604}
1605
1606/// Returns the innermost common loop depth for the set of operations in 'ops'.
1607// TODO: Move this to LoopUtils.
1609 ArrayRef<Operation *> ops, SmallVectorImpl<AffineForOp> *surroundingLoops) {
1610 unsigned numOps = ops.size();
1611 assert(numOps > 0 && "Expected at least one operation");
1612
1613 std::vector<SmallVector<AffineForOp, 4>> loops(numOps);
1614 unsigned loopDepthLimit = std::numeric_limits<unsigned>::max();
1615 for (unsigned i = 0; i < numOps; ++i) {
1616 getAffineForIVs(*ops[i], &loops[i]);
1617 loopDepthLimit = std::min(loopDepthLimit, (unsigned)loops[i].size());
1618 }
1619
1620 unsigned loopDepth = 0;
1621 for (unsigned d = 0; d < loopDepthLimit; ++d) {
1622 unsigned i;
1623 for (i = 1; i < numOps; ++i) {
1624 if (loops[i - 1][d] != loops[i][d])
1625 return loopDepth;
1626 }
1627 if (surroundingLoops)
1628 surroundingLoops->push_back(loops[i - 1][d]);
1629 ++loopDepth;
1630 }
1631 return loopDepth;
1632}
1633
1634/// Computes in 'sliceUnion' the union of all slice bounds computed at
1635/// 'loopDepth' between all dependent pairs of ops in 'opsA' and 'opsB', and
1636/// then verifies if it is valid. Returns 'SliceComputationResult::Success' if
1637/// union was computed correctly, an appropriate failure otherwise.
1640 ArrayRef<Operation *> opsB, unsigned loopDepth,
1641 unsigned numCommonLoops, bool isBackwardSlice,
1642 ComputationSliceState *sliceUnion) {
1643 // Compute the union of slice bounds between all pairs in 'opsA' and
1644 // 'opsB' in 'sliceUnionCst'.
1645 FlatAffineValueConstraints sliceUnionCst;
1646 assert(sliceUnionCst.getNumDimAndSymbolVars() == 0);
1647 std::vector<std::pair<Operation *, Operation *>> dependentOpPairs;
1648 MemRefAccess srcAccess;
1649 MemRefAccess dstAccess;
1650 for (Operation *a : opsA) {
1651 srcAccess = MemRefAccess(a);
1652 for (Operation *b : opsB) {
1653 dstAccess = MemRefAccess(b);
1654 if (srcAccess.memref != dstAccess.memref)
1655 continue;
1656 // Check if 'loopDepth' exceeds nesting depth of src/dst ops.
1657 if ((!isBackwardSlice && loopDepth > getNestingDepth(a)) ||
1658 (isBackwardSlice && loopDepth > getNestingDepth(b))) {
1659 LDBG() << "Invalid loop depth";
1661 }
1662
1663 bool readReadAccesses = isa<AffineReadOpInterface>(srcAccess.opInst) &&
1664 isa<AffineReadOpInterface>(dstAccess.opInst);
1665 FlatAffineValueConstraints dependenceConstraints;
1666 // Check dependence between 'srcAccess' and 'dstAccess'.
1668 srcAccess, dstAccess, /*loopDepth=*/numCommonLoops + 1,
1669 &dependenceConstraints, /*dependenceComponents=*/nullptr,
1670 /*allowRAR=*/readReadAccesses);
1671 if (result.value == DependenceResult::Failure) {
1672 LDBG() << "Dependence check failed";
1674 }
1676 continue;
1677 dependentOpPairs.emplace_back(a, b);
1678
1679 // Compute slice bounds for 'srcAccess' and 'dstAccess'.
1680 ComputationSliceState tmpSliceState;
1681 getComputationSliceState(a, b, dependenceConstraints, loopDepth,
1682 isBackwardSlice, &tmpSliceState);
1683
1684 if (sliceUnionCst.getNumDimAndSymbolVars() == 0) {
1685 // Initialize 'sliceUnionCst' with the bounds computed in previous step.
1686 if (failed(tmpSliceState.getAsConstraints(&sliceUnionCst))) {
1687 LDBG() << "Unable to compute slice bound constraints";
1689 }
1690 assert(sliceUnionCst.getNumDimAndSymbolVars() > 0);
1691 continue;
1692 }
1693
1694 // Compute constraints for 'tmpSliceState' in 'tmpSliceCst'.
1695 FlatAffineValueConstraints tmpSliceCst;
1696 if (failed(tmpSliceState.getAsConstraints(&tmpSliceCst))) {
1697 LDBG() << "Unable to compute slice bound constraints";
1699 }
1700
1701 // Align coordinate spaces of 'sliceUnionCst' and 'tmpSliceCst' if needed.
1702 if (!sliceUnionCst.areVarsAlignedWithOther(tmpSliceCst)) {
1703
1704 // Pre-constraint var alignment: record loop IVs used in each constraint
1705 // system.
1706 SmallPtrSet<Value, 8> sliceUnionIVs;
1707 for (unsigned k = 0, l = sliceUnionCst.getNumDimVars(); k < l; ++k)
1708 sliceUnionIVs.insert(sliceUnionCst.getValue(k));
1709 SmallPtrSet<Value, 8> tmpSliceIVs;
1710 for (unsigned k = 0, l = tmpSliceCst.getNumDimVars(); k < l; ++k)
1711 tmpSliceIVs.insert(tmpSliceCst.getValue(k));
1712
1713 sliceUnionCst.mergeAndAlignVarsWithOther(/*offset=*/0, &tmpSliceCst);
1714
1715 // Post-constraint var alignment: add loop IV bounds missing after
1716 // var alignment to constraint systems. This can occur if one constraint
1717 // system uses an loop IV that is not used by the other. The call
1718 // to unionBoundingBox below expects constraints for each Loop IV, even
1719 // if they are the unsliced full loop bounds added here.
1720 if (failed(addMissingLoopIVBounds(sliceUnionIVs, &sliceUnionCst)))
1722 if (failed(addMissingLoopIVBounds(tmpSliceIVs, &tmpSliceCst)))
1724 }
1725 // Compute union bounding box of 'sliceUnionCst' and 'tmpSliceCst'.
1726 if (sliceUnionCst.getNumLocalVars() > 0 ||
1727 tmpSliceCst.getNumLocalVars() > 0 ||
1728 failed(sliceUnionCst.unionBoundingBox(tmpSliceCst))) {
1729 LDBG() << "Unable to compute union bounding box of slice bounds";
1731 }
1732 }
1733 }
1734
1735 // Empty union.
1736 if (sliceUnionCst.getNumDimAndSymbolVars() == 0) {
1737 LDBG() << "empty slice union - unexpected";
1739 }
1740
1741 // Gather loops surrounding ops from loop nest where slice will be inserted.
1743 for (auto &dep : dependentOpPairs) {
1744 ops.push_back(isBackwardSlice ? dep.second : dep.first);
1745 }
1746 SmallVector<AffineForOp, 4> surroundingLoops;
1747 unsigned innermostCommonLoopDepth =
1748 getInnermostCommonLoopDepth(ops, &surroundingLoops);
1749 if (loopDepth > innermostCommonLoopDepth) {
1750 LDBG() << "Exceeds max loop depth";
1752 }
1753
1754 // Store 'numSliceLoopIVs' before converting dst loop IVs to dims.
1755 unsigned numSliceLoopIVs = sliceUnionCst.getNumDimVars();
1756
1757 // Convert any dst loop IVs which are symbol variables to dim variables.
1758 sliceUnionCst.convertLoopIVSymbolsToDims();
1759 sliceUnion->clearBounds();
1760 sliceUnion->lbs.resize(numSliceLoopIVs, AffineMap());
1761 sliceUnion->ubs.resize(numSliceLoopIVs, AffineMap());
1762
1763 // Get slice bounds from slice union constraints 'sliceUnionCst'.
1764 sliceUnionCst.getSliceBounds(/*offset=*/0, numSliceLoopIVs,
1765 opsA[0]->getContext(), &sliceUnion->lbs,
1766 &sliceUnion->ubs, /*closedUb=*/false,
1767 /*allowMultiResultUb=*/true);
1768
1769 // Add slice bound operands of union.
1770 SmallVector<Value, 4> sliceBoundOperands;
1771 sliceUnionCst.getValues(numSliceLoopIVs,
1772 sliceUnionCst.getNumDimAndSymbolVars(),
1773 &sliceBoundOperands);
1774
1775 // Copy src loop IVs from 'sliceUnionCst' to 'sliceUnion'.
1776 sliceUnion->ivs.clear();
1777 sliceUnionCst.getValues(0, numSliceLoopIVs, &sliceUnion->ivs);
1778
1779 // Set loop nest insertion point to block start at 'loopDepth' for forward
1780 // slices, while at the end for backward slices.
1781 sliceUnion->insertPoint =
1782 isBackwardSlice
1783 ? surroundingLoops[loopDepth - 1].getBody()->begin()
1784 : std::prev(surroundingLoops[loopDepth - 1].getBody()->end());
1785
1786 // Give each bound its own copy of 'sliceBoundOperands' for subsequent
1787 // canonicalization.
1788 sliceUnion->lbOperands.resize(numSliceLoopIVs, sliceBoundOperands);
1789 sliceUnion->ubOperands.resize(numSliceLoopIVs, sliceBoundOperands);
1790
1791 // Check if the slice computed is valid. Return success only if it is verified
1792 // that the slice is valid, otherwise return appropriate failure status.
1793 std::optional<bool> isSliceValid = sliceUnion->isSliceValid();
1794 if (!isSliceValid) {
1795 LDBG() << "Cannot determine if the slice is valid";
1797 }
1798 if (!*isSliceValid)
1800
1802}
1803
1804/// Returns the number of iterations the slice bounded below by `lbMap` and
1805/// above by `ubMap` runs for, where that is a constant.
1806///
1807/// An upper bound of several results is the min of them, so each result taken
1808/// against the lower bound bounds the count from above and the smallest of
1809/// those that comes out constant is the tightest constant bound there is. A
1810/// tiled loop clamped at the end of the data has exactly this shape --
1811/// `min(%i * 64 + 64, 1000)` over `%i * 64` -- where the tile-relative result
1812/// gives the 64 and the extent gives nothing constant at all.
1813static std::optional<uint64_t> getConstDifference(AffineMap lbMap,
1814 AffineMap ubMap) {
1815 assert(lbMap.getNumResults() == 1 && "expected single result lower bound");
1816 assert(ubMap.getNumResults() >= 1 && "expected at least one upper bound");
1817 assert(lbMap.getNumDims() == ubMap.getNumDims());
1818 assert(lbMap.getNumSymbols() == ubMap.getNumSymbols());
1819 AffineExpr lbExpr(lbMap.getResult(0));
1820 std::optional<uint64_t> tripCount;
1821 for (AffineExpr ubExpr : ubMap.getResults()) {
1822 AffineExpr loopSpanExpr = simplifyAffineExpr(
1823 ubExpr - lbExpr, lbMap.getNumDims(), lbMap.getNumSymbols());
1824 auto cExpr = dyn_cast<AffineConstantExpr>(loopSpanExpr);
1825 if (!cExpr)
1826 continue;
1827 if (cExpr.getValue() < 0)
1828 return 0;
1829 tripCount =
1830 std::min(tripCount.value_or(std::numeric_limits<uint64_t>::max()),
1831 (uint64_t)cExpr.getValue());
1832 }
1833 return tripCount;
1834}
1835
1836// Builds a map 'tripCountMap' from AffineForOp to constant trip count for loop
1837// nest surrounding represented by slice loop bounds in 'slice'. Returns true
1838// on success, false otherwise (if a non-constant trip count was encountered).
1839// TODO: Make this work with non-unit step loops.
1841 const ComputationSliceState &slice,
1842 llvm::SmallDenseMap<Operation *, uint64_t, 8> *tripCountMap) {
1843 unsigned numSrcLoopIVs = slice.ivs.size();
1844 // Populate map from AffineForOp -> trip count
1845 for (unsigned i = 0; i < numSrcLoopIVs; ++i) {
1846 AffineForOp forOp = getForInductionVarOwner(slice.ivs[i]);
1847 auto *op = forOp.getOperation();
1848 AffineMap lbMap = slice.lbs[i];
1849 AffineMap ubMap = slice.ubs[i];
1850 // If lower or upper bound maps are null or provide no results, it implies
1851 // that source loop was not at all sliced, and the entire loop will be a
1852 // part of the slice.
1853 if (!lbMap || lbMap.getNumResults() == 0 || !ubMap ||
1854 ubMap.getNumResults() == 0) {
1855 // The iteration of src loop IV 'i' was not sliced. Use full loop bounds.
1856 if (forOp.hasConstantLowerBound() && forOp.hasConstantUpperBound()) {
1857 (*tripCountMap)[op] =
1858 forOp.getConstantUpperBound() - forOp.getConstantLowerBound();
1859 continue;
1860 }
1861 std::optional<APInt> maybeConstTripCount = forOp.getStaticTripCount();
1862 if (maybeConstTripCount.has_value()) {
1863 (*tripCountMap)[op] = maybeConstTripCount->getZExtValue();
1864 continue;
1865 }
1866 return false;
1867 }
1868 std::optional<uint64_t> tripCount = getConstDifference(lbMap, ubMap);
1869 // Slice bounds are created with a constant ub - lb difference.
1870 if (!tripCount.has_value())
1871 return false;
1872 (*tripCountMap)[op] = *tripCount;
1873 }
1874 return true;
1875}
1876
1877// Return the number of iterations in the given slice.
1879 const llvm::SmallDenseMap<Operation *, uint64_t, 8> &sliceTripCountMap) {
1880 uint64_t iterCount = 1;
1881 for (const auto &count : sliceTripCountMap) {
1882 iterCount *= count.second;
1883 }
1884 return iterCount;
1885}
1886
1887const char *const kSliceFusionBarrierAttrName = "slice_fusion_barrier";
1888// Computes slice bounds by projecting out any loop IVs from
1889// 'dependenceConstraints' at depth greater than 'loopDepth', and computes slice
1890// bounds in 'sliceState' which represent the one loop nest's IVs in terms of
1891// the other loop nest's IVs, symbols and constants (using 'isBackwardsSlice').
1893 Operation *depSourceOp, Operation *depSinkOp,
1894 const FlatAffineValueConstraints &dependenceConstraints, unsigned loopDepth,
1895 bool isBackwardSlice, ComputationSliceState *sliceState) {
1896 // Get loop nest surrounding src operation.
1897 SmallVector<AffineForOp, 4> srcLoopIVs;
1898 getAffineForIVs(*depSourceOp, &srcLoopIVs);
1899 unsigned numSrcLoopIVs = srcLoopIVs.size();
1900
1901 // Get loop nest surrounding dst operation.
1902 SmallVector<AffineForOp, 4> dstLoopIVs;
1903 getAffineForIVs(*depSinkOp, &dstLoopIVs);
1904 unsigned numDstLoopIVs = dstLoopIVs.size();
1905
1906 assert((!isBackwardSlice && loopDepth <= numSrcLoopIVs) ||
1907 (isBackwardSlice && loopDepth <= numDstLoopIVs));
1908
1909 // Project out dimensions other than those up to 'loopDepth'.
1910 unsigned pos = isBackwardSlice ? numSrcLoopIVs + loopDepth : loopDepth;
1911 unsigned num =
1912 isBackwardSlice ? numDstLoopIVs - loopDepth : numSrcLoopIVs - loopDepth;
1913 FlatAffineValueConstraints sliceCst(dependenceConstraints);
1914 sliceCst.projectOut(pos, num);
1915
1916 // Add slice loop IV values to 'sliceState'.
1917 unsigned offset = isBackwardSlice ? 0 : loopDepth;
1918 unsigned numSliceLoopIVs = isBackwardSlice ? numSrcLoopIVs : numDstLoopIVs;
1919 sliceCst.getValues(offset, offset + numSliceLoopIVs, &sliceState->ivs);
1920
1921 // Set up lower/upper bound affine maps for the slice.
1922 sliceState->lbs.resize(numSliceLoopIVs, AffineMap());
1923 sliceState->ubs.resize(numSliceLoopIVs, AffineMap());
1924
1925 // Get bounds for slice IVs in terms of other IVs, symbols, and constants.
1926 sliceCst.getSliceBounds(offset, numSliceLoopIVs, depSourceOp->getContext(),
1927 &sliceState->lbs, &sliceState->ubs,
1928 /*closedUb=*/false, /*allowMultiResultUb=*/true);
1929
1930 // Set up bound operands for the slice's lower and upper bounds.
1931 SmallVector<Value, 4> sliceBoundOperands;
1932 unsigned numDimsAndSymbols = sliceCst.getNumDimAndSymbolVars();
1933 for (unsigned i = 0; i < numDimsAndSymbols; ++i) {
1934 if (i < offset || i >= offset + numSliceLoopIVs)
1935 sliceBoundOperands.push_back(sliceCst.getValue(i));
1936 }
1937
1938 // Give each bound its own copy of 'sliceBoundOperands' for subsequent
1939 // canonicalization.
1940 sliceState->lbOperands.resize(numSliceLoopIVs, sliceBoundOperands);
1941 sliceState->ubOperands.resize(numSliceLoopIVs, sliceBoundOperands);
1942
1943 // Set destination loop nest insertion point to block start at 'dstLoopDepth'.
1944 sliceState->insertPoint =
1945 isBackwardSlice ? dstLoopIVs[loopDepth - 1].getBody()->begin()
1946 : std::prev(srcLoopIVs[loopDepth - 1].getBody()->end());
1947
1948 llvm::SmallDenseSet<Value, 8> sequentialLoops;
1949 if (isa<AffineReadOpInterface>(depSourceOp) &&
1950 isa<AffineReadOpInterface>(depSinkOp)) {
1951 // For read-read access pairs, clear any slice bounds on sequential loops.
1952 // Get sequential loops in loop nest rooted at 'srcLoopIVs[0]'.
1953 getSequentialLoops(isBackwardSlice ? srcLoopIVs[0] : dstLoopIVs[0],
1954 &sequentialLoops);
1955 }
1956 auto getSliceLoop = [&](unsigned i) {
1957 return isBackwardSlice ? srcLoopIVs[i] : dstLoopIVs[i];
1958 };
1959 auto isInnermostInsertion = [&]() {
1960 return (isBackwardSlice ? loopDepth >= srcLoopIVs.size()
1961 : loopDepth >= dstLoopIVs.size());
1962 };
1963 llvm::SmallDenseMap<Operation *, uint64_t, 8> sliceTripCountMap;
1964 auto srcIsUnitSlice = [&]() {
1965 return (buildSliceTripCountMap(*sliceState, &sliceTripCountMap) &&
1966 (getSliceIterationCount(sliceTripCountMap) == 1));
1967 };
1968 // Clear all sliced loop bounds beginning at the first sequential loop, or
1969 // first loop with a slice fusion barrier attribute..
1970
1971 for (unsigned i = 0; i < numSliceLoopIVs; ++i) {
1972 Value iv = getSliceLoop(i).getInductionVar();
1973 if (sequentialLoops.count(iv) == 0 &&
1974 getSliceLoop(i)->getDiscardableAttr(kSliceFusionBarrierAttrName) ==
1975 nullptr)
1976 continue;
1977 // Skip reset of bounds of reduction loop inserted in the destination loop
1978 // that meets the following conditions:
1979 // 1. Slice is single trip count.
1980 // 2. Loop bounds of the source and destination match.
1981 // 3. Is being inserted at the innermost insertion point.
1982 std::optional<bool> isMaximal = sliceState->isMaximal();
1983 if (isLoopParallelAndContainsReduction(getSliceLoop(i)) &&
1984 isInnermostInsertion() && srcIsUnitSlice() && isMaximal && *isMaximal)
1985 continue;
1986 for (unsigned j = i; j < numSliceLoopIVs; ++j) {
1987 sliceState->lbs[j] = AffineMap();
1988 sliceState->ubs[j] = AffineMap();
1989 }
1990 break;
1991 }
1992}
1993
1994/// Creates a computation slice of the loop nest surrounding 'srcOpInst',
1995/// updates the slice loop bounds with any non-null bound maps specified in
1996/// 'sliceState', and inserts this slice into the loop nest surrounding
1997/// 'dstOpInst' at loop depth 'dstLoopDepth'.
1998// TODO: extend the slicing utility to compute slices that
1999// aren't necessarily a one-to-one relation b/w the source and destination. The
2000// relation between the source and destination could be many-to-many in general.
2001// TODO: the slice computation is incorrect in the cases
2002// where the dependence from the source to the destination does not cover the
2003// entire destination index set. Subtract out the dependent destination
2004// iterations from destination index set and check for emptiness --- this is one
2005// solution.
2007 Operation *srcOpInst, Operation *dstOpInst, unsigned dstLoopDepth,
2008 ComputationSliceState *sliceState) {
2009 // Get loop nest surrounding src operation.
2010 SmallVector<AffineForOp, 4> srcLoopIVs;
2011 getAffineForIVs(*srcOpInst, &srcLoopIVs);
2012 unsigned numSrcLoopIVs = srcLoopIVs.size();
2013
2014 // Get loop nest surrounding dst operation.
2015 SmallVector<AffineForOp, 4> dstLoopIVs;
2016 getAffineForIVs(*dstOpInst, &dstLoopIVs);
2017 unsigned dstLoopIVsSize = dstLoopIVs.size();
2018 if (dstLoopDepth > dstLoopIVsSize) {
2019 dstOpInst->emitError("invalid destination loop depth");
2020 return AffineForOp();
2021 }
2022
2023 // Find the op block positions of 'srcOpInst' within 'srcLoopIVs'.
2024 SmallVector<unsigned, 4> positions;
2025 // TODO: This code is incorrect since srcLoopIVs can be 0-d.
2026 findInstPosition(srcOpInst, srcLoopIVs[0]->getBlock(), &positions);
2027
2028 // Clone src loop nest and insert it a the beginning of the operation block
2029 // of the loop at 'dstLoopDepth' in 'dstLoopIVs'.
2030 auto dstAffineForOp = dstLoopIVs[dstLoopDepth - 1];
2031 OpBuilder b(dstAffineForOp.getBody(), dstAffineForOp.getBody()->begin());
2032 auto sliceLoopNest =
2033 cast<AffineForOp>(b.clone(*srcLoopIVs[0].getOperation()));
2034
2035 Operation *sliceInst =
2036 getInstAtPosition(positions, /*level=*/0, sliceLoopNest.getBody());
2037 // Get loop nest surrounding 'sliceInst'.
2038 SmallVector<AffineForOp, 4> sliceSurroundingLoops;
2039 getAffineForIVs(*sliceInst, &sliceSurroundingLoops);
2040
2041 // Sanity check.
2042 unsigned sliceSurroundingLoopsSize = sliceSurroundingLoops.size();
2043 (void)sliceSurroundingLoopsSize;
2044 assert(dstLoopDepth + numSrcLoopIVs >= sliceSurroundingLoopsSize);
2045 unsigned sliceLoopLimit = dstLoopDepth + numSrcLoopIVs;
2046 (void)sliceLoopLimit;
2047 assert(sliceLoopLimit >= sliceSurroundingLoopsSize);
2048
2049 // Update loop bounds for loops in 'sliceLoopNest'.
2050 for (unsigned i = 0; i < numSrcLoopIVs; ++i) {
2051 auto forOp = sliceSurroundingLoops[dstLoopDepth + i];
2052 if (AffineMap lbMap = sliceState->lbs[i])
2053 forOp.setLowerBound(sliceState->lbOperands[i], lbMap);
2054 if (AffineMap ubMap = sliceState->ubs[i])
2055 forOp.setUpperBound(sliceState->ubOperands[i], ubMap);
2056 }
2057 return sliceLoopNest;
2058}
2059
2060// Constructs MemRefAccess populating it with the memref, its indices and
2061// opinst from 'loadOrStoreOpInst'.
2063 if (auto loadOp = dyn_cast<AffineReadOpInterface>(memOp)) {
2064 memref = loadOp.getMemRef();
2065 opInst = memOp;
2066 llvm::append_range(indices, loadOp.getMapOperands());
2067 } else {
2068 assert(isa<AffineWriteOpInterface>(memOp) &&
2069 "Affine read/write op expected");
2070 auto storeOp = cast<AffineWriteOpInterface>(memOp);
2071 opInst = memOp;
2072 memref = storeOp.getMemRef();
2073 llvm::append_range(indices, storeOp.getMapOperands());
2074 }
2075}
2076
2077unsigned MemRefAccess::getRank() const {
2078 return cast<MemRefType>(memref.getType()).getRank();
2079}
2080
2082 return isa<AffineWriteOpInterface>(opInst);
2083}
2084
2085/// Returns the nesting depth of this statement, i.e., the number of loops
2086/// surrounding this statement.
2088 Operation *currOp = op;
2089 unsigned depth = 0;
2090 while ((currOp = currOp->getParentOp())) {
2091 if (isa<AffineForOp>(currOp))
2092 depth++;
2093 if (auto parOp = dyn_cast<AffineParallelOp>(currOp))
2094 depth += parOp.getNumDims();
2095 }
2096 return depth;
2097}
2098
2099/// Equal if both affine accesses are provably equivalent (at compile
2100/// time) when considering the memref, the affine maps and their respective
2101/// operands. The equality of access functions + operands is checked by
2102/// subtracting fully composed value maps, and then simplifying the difference
2103/// using the expression flattener.
2104/// TODO: this does not account for aliasing of memrefs.
2106 if (memref != rhs.memref)
2107 return false;
2108
2109 AffineValueMap diff, thisMap, rhsMap;
2110 getAccessMap(&thisMap);
2111 rhs.getAccessMap(&rhsMap);
2112 return thisMap == rhsMap;
2113}
2114
2116 auto *currOp = op.getParentOp();
2117 AffineForOp currAffineForOp;
2118 // Traverse up the hierarchy collecting all 'affine.for' and affine.parallel
2119 // operation while skipping over 'affine.if' operations.
2120 while (currOp) {
2121 if (AffineForOp currAffineForOp = dyn_cast<AffineForOp>(currOp))
2122 ivs.push_back(currAffineForOp.getInductionVar());
2123 else if (auto parOp = dyn_cast<AffineParallelOp>(currOp))
2124 llvm::append_range(ivs, parOp.getIVs());
2125 currOp = currOp->getParentOp();
2126 }
2127 std::reverse(ivs.begin(), ivs.end());
2128}
2129
2130/// Returns the number of surrounding loops common to 'loopsA' and 'loopsB',
2131/// where each lists loops from outer-most to inner-most in loop nest.
2133 Operation &b) {
2134 SmallVector<Value, 4> loopsA, loopsB;
2135 getAffineIVs(a, loopsA);
2136 getAffineIVs(b, loopsB);
2137
2138 unsigned minNumLoops = std::min(loopsA.size(), loopsB.size());
2139 unsigned numCommonLoops = 0;
2140 for (unsigned i = 0; i < minNumLoops; ++i) {
2141 if (loopsA[i] != loopsB[i])
2142 break;
2143 ++numCommonLoops;
2144 }
2145 return numCommonLoops;
2146}
2147
2148static std::optional<int64_t> getMemoryFootprintBytes(Block &block,
2149 Block::iterator start,
2150 Block::iterator end,
2151 int memorySpace) {
2153
2154 // Walk this 'affine.for' operation to gather all memory regions.
2155 auto result = block.walk(start, end, [&](Operation *opInst) -> WalkResult {
2156 if (!isa<AffineReadOpInterface, AffineWriteOpInterface>(opInst)) {
2157 // Neither load nor a store op.
2158 return WalkResult::advance();
2159 }
2160
2161 // Compute the memref region symbolic in any IVs enclosing this block.
2162 auto region = std::make_unique<MemRefRegion>(opInst->getLoc());
2163 if (failed(
2164 region->compute(opInst,
2165 /*loopDepth=*/getNestingDepth(&*block.begin())))) {
2166 LDBG() << "Error obtaining memory region";
2167 opInst->emitError("error obtaining memory region");
2168 return failure();
2169 }
2170
2171 auto [it, inserted] = regions.try_emplace(region->memref);
2172 if (inserted) {
2173 it->second = std::move(region);
2174 } else if (failed(it->second->unionBoundingBox(*region))) {
2175 LDBG() << "getMemoryFootprintBytes: unable to perform a union on a "
2176 "memory region";
2177 opInst->emitWarning(
2178 "getMemoryFootprintBytes: unable to perform a union on a memory "
2179 "region");
2180 return failure();
2181 }
2182 return WalkResult::advance();
2183 });
2184 if (result.wasInterrupted())
2185 return std::nullopt;
2186
2187 int64_t totalSizeInBytes = 0;
2188 for (const auto &region : regions) {
2189 std::optional<int64_t> size = region.second->getRegionSize();
2190 if (!size.has_value())
2191 return std::nullopt;
2192 totalSizeInBytes += *size;
2193 }
2194 return totalSizeInBytes;
2195}
2196
2197std::optional<int64_t> mlir::affine::getMemoryFootprintBytes(AffineForOp forOp,
2198 int memorySpace) {
2199 auto *forInst = forOp.getOperation();
2200 return ::getMemoryFootprintBytes(
2201 *forInst->getBlock(), Block::iterator(forInst),
2202 std::next(Block::iterator(forInst)), memorySpace);
2203}
2204
2205/// Returns whether a loop is parallel and contains a reduction loop.
2207 SmallVector<LoopReduction> reductions;
2208 if (!isLoopParallel(forOp, &reductions))
2209 return false;
2210 return !reductions.empty();
2211}
2212
2213/// Returns in 'sequentialLoops' all sequential loops in loop nest rooted
2214/// at 'forOp'.
2216 AffineForOp forOp, llvm::SmallDenseSet<Value, 8> *sequentialLoops) {
2217 forOp->walk([&](Operation *op) {
2218 if (auto innerFor = dyn_cast<AffineForOp>(op))
2219 if (!isLoopParallel(innerFor))
2220 sequentialLoops->insert(innerFor.getInductionVar());
2221 });
2222}
2223
2225 FailureOr<FlatAffineValueConstraints> fac =
2227 // Semi-affine sets can't be flattened; return them as is.
2228 if (failed(fac))
2229 return set;
2230 if (fac->isEmpty())
2232 set.getContext());
2233 fac->removeTrivialRedundancy();
2234
2235 auto simplifiedSet = fac->getAsIntegerSet(set.getContext());
2236 assert(simplifiedSet && "guaranteed to succeed while roundtripping");
2237 return simplifiedSet;
2238}
2239
2240static void unpackOptionalValues(ArrayRef<std::optional<Value>> source,
2242 target = llvm::map_to_vector<4>(source, [](std::optional<Value> val) {
2243 return val.has_value() ? *val : Value();
2244 });
2245}
2246
2247/// Bound an identifier `pos` in a given FlatAffineValueConstraints with
2248/// constraints drawn from an affine map. Before adding the constraint, the
2249/// dimensions/symbols of the affine map are aligned with `constraints`.
2250/// `operands` are the SSA Value operands used with the affine map.
2251/// Note: This function adds a new symbol column to the `constraints` for each
2252/// dimension/symbol that exists in the affine map but not in `constraints`.
2253static LogicalResult alignAndAddBound(FlatAffineValueConstraints &constraints,
2254 BoundType type, unsigned pos,
2255 AffineMap map, ValueRange operands) {
2256 SmallVector<Value> dims, syms, newSyms;
2257 unpackOptionalValues(constraints.getMaybeValues(VarKind::SetDim), dims);
2258 unpackOptionalValues(constraints.getMaybeValues(VarKind::Symbol), syms);
2259
2260 AffineMap alignedMap =
2261 alignAffineMapWithValues(map, operands, dims, syms, &newSyms);
2262 for (unsigned i = syms.size(); i < newSyms.size(); ++i)
2263 constraints.appendSymbolVar(newSyms[i]);
2264 return constraints.addBound(type, pos, alignedMap);
2265}
2266
2267/// Add `val` to each result of `map`.
2269 SmallVector<AffineExpr> newResults;
2270 for (AffineExpr r : map.getResults())
2271 newResults.push_back(r + val);
2272 return AffineMap::get(map.getNumDims(), map.getNumSymbols(), newResults,
2273 map.getContext());
2274}
2275
2276// Attempt to simplify the given min/max operation by proving that its value is
2277// bounded by the same lower and upper bound.
2278//
2279// Bounds are computed by FlatAffineValueConstraints. Invariants required for
2280// finding/proving bounds should be supplied via `constraints`.
2281//
2282// 1. Add dimensions for `op` and `opBound` (lower or upper bound of `op`).
2283// 2. Compute an upper bound of `op` (in case of `isMin`) or a lower bound (in
2284// case of `!isMin`) and bind it to `opBound`. SSA values that are used in
2285// `op` but are not part of `constraints`, are added as extra symbols.
2286// 3. For each result of `op`: Add result as a dimension `r_i`. Prove that:
2287// * If `isMin`: r_i >= opBound
2288// * If `isMax`: r_i <= opBound
2289// If this is the case, ub(op) == lb(op).
2290// 4. Replace `op` with `opBound`.
2291//
2292// In summary, the following constraints are added throughout this function.
2293// Note: `invar` are dimensions added by the caller to express the invariants.
2294// (Showing only the case where `isMin`.)
2295//
2296// invar | op | opBound | r_i | extra syms... | const | eq/ineq
2297// ------+-------+---------+-----+---------------+-------+-------------------
2298// (various eq./ineq. constraining `invar`, added by the caller)
2299// ... | 0 | 0 | 0 | 0 | ... | ...
2300// ------+-------+---------+-----+---------------+-------+-------------------
2301// (various ineq. constraining `op` in terms of `op` operands (`invar` and
2302// extra `op` operands "extra syms" that are not in `invar`)).
2303// ... | -1 | 0 | 0 | ... | ... | >= 0
2304// ------+-------+---------+-----+---------------+-------+-------------------
2305// (set `opBound` to `op` upper bound in terms of `invar` and "extra syms")
2306// ... | 0 | -1 | 0 | ... | ... | = 0
2307// ------+-------+---------+-----+---------------+-------+-------------------
2308// (for each `op` map result r_i: set r_i to corresponding map result,
2309// prove that r_i >= minOpUb via contradiction)
2310// ... | 0 | 0 | -1 | ... | ... | = 0
2311// 0 | 0 | 1 | -1 | 0 | -1 | >= 0
2312//
2314 Operation *op, FlatAffineValueConstraints constraints) {
2315 bool isMin = isa<AffineMinOp>(op);
2316 assert((isMin || isa<AffineMaxOp>(op)) && "expect AffineMin/MaxOp");
2317 MLIRContext *ctx = op->getContext();
2318 Builder builder(ctx);
2319 AffineMap map =
2320 isMin ? cast<AffineMinOp>(op).getMap() : cast<AffineMaxOp>(op).getMap();
2321 ValueRange operands = op->getOperands();
2322 unsigned numResults = map.getNumResults();
2323
2324 // Add a few extra dimensions.
2325 unsigned dimOp = constraints.appendDimVar(); // `op`
2326 unsigned dimOpBound = constraints.appendDimVar(); // `op` lower/upper bound
2327 unsigned resultDimStart = constraints.appendDimVar(/*num=*/numResults);
2328
2329 // Add an inequality for each result expr_i of map:
2330 // isMin: op <= expr_i, !isMin: op >= expr_i
2331 auto boundType = isMin ? BoundType::UB : BoundType::LB;
2332 // Upper bounds are exclusive, so add 1. (`affine.min` ops are inclusive.)
2333 AffineMap mapLbUb = isMin ? addConstToResults(map, 1) : map;
2334 if (failed(
2335 alignAndAddBound(constraints, boundType, dimOp, mapLbUb, operands)))
2336 return failure();
2337
2338 // Try to compute a lower/upper bound for op, expressed in terms of the other
2339 // `dims` and extra symbols.
2340 SmallVector<AffineMap> opLb(1), opUb(1);
2341 constraints.getSliceBounds(dimOp, 1, ctx, &opLb, &opUb);
2342 AffineMap sliceBound = isMin ? opUb[0] : opLb[0];
2343 // TODO: `getSliceBounds` may return multiple bounds at the moment. This is
2344 // a TODO of `getSliceBounds` and not handled here.
2345 if (!sliceBound || sliceBound.getNumResults() != 1)
2346 return failure(); // No or multiple bounds found.
2347 // Recover the inclusive UB in the case of an `affine.min`.
2348 AffineMap boundMap = isMin ? addConstToResults(sliceBound, -1) : sliceBound;
2349
2350 // Add an equality: Set dimOpBound to computed bound.
2351 // Add back dimension for op. (Was removed by `getSliceBounds`.)
2352 AffineMap alignedBoundMap = boundMap.shiftDims(/*shift=*/1, /*offset=*/dimOp);
2353 if (failed(constraints.addBound(BoundType::EQ, dimOpBound, alignedBoundMap)))
2354 return failure();
2355
2356 // If the constraint system is empty, there is an inconsistency. (E.g., this
2357 // can happen if loop lb > ub.)
2358 if (constraints.isEmpty())
2359 return failure();
2360
2361 // In the case of `isMin` (`!isMin` is inversed):
2362 // Prove that each result of `map` has a lower bound that is equal to (or
2363 // greater than) the upper bound of `op` (`dimOpBound`). In that case, `op`
2364 // can be replaced with the bound. I.e., prove that for each result
2365 // expr_i (represented by dimension r_i):
2366 //
2367 // r_i >= opBound
2368 //
2369 // To prove this inequality, add its negation to the constraint set and prove
2370 // that the constraint set is empty.
2371 for (unsigned i = resultDimStart; i < resultDimStart + numResults; ++i) {
2372 FlatAffineValueConstraints newConstr(constraints);
2373
2374 // Add an equality: r_i = expr_i
2375 // Note: These equalities could have been added earlier and used to express
2376 // minOp <= expr_i. However, then we run the risk that `getSliceBounds`
2377 // computes minOpUb in terms of r_i dims, which is not desired.
2378 if (failed(alignAndAddBound(newConstr, BoundType::EQ, i,
2379 map.getSubMap({i - resultDimStart}), operands)))
2380 return failure();
2381
2382 // If `isMin`: Add inequality: r_i < opBound
2383 // equiv.: opBound - r_i - 1 >= 0
2384 // If `!isMin`: Add inequality: r_i > opBound
2385 // equiv.: -opBound + r_i - 1 >= 0
2386 SmallVector<int64_t> ineq(newConstr.getNumCols(), 0);
2387 ineq[dimOpBound] = isMin ? 1 : -1;
2388 ineq[i] = isMin ? -1 : 1;
2389 ineq[newConstr.getNumCols() - 1] = -1;
2390 newConstr.addInequality(ineq);
2391 if (!newConstr.isEmpty())
2392 return failure();
2393 }
2394
2395 // Lower and upper bound of `op` are equal. Replace `minOp` with its bound.
2396 AffineMap newMap = alignedBoundMap;
2397 SmallVector<Value> newOperands;
2398 unpackOptionalValues(constraints.getMaybeValues(), newOperands);
2399 // If dims/symbols have known constant values, use those in order to simplify
2400 // the affine map further.
2401 for (int64_t i = 0, e = constraints.getNumDimAndSymbolVars(); i < e; ++i) {
2402 // Skip unused operands and operands that are already constants.
2403 if (!newOperands[i] || getConstantIntValue(newOperands[i]))
2404 continue;
2405 if (auto bound = constraints.getConstantBound64(BoundType::EQ, i)) {
2406 AffineExpr expr =
2407 i < newMap.getNumDims()
2408 ? builder.getAffineDimExpr(i)
2409 : builder.getAffineSymbolExpr(i - newMap.getNumDims());
2410 newMap = newMap.replace(expr, builder.getAffineConstantExpr(*bound),
2411 newMap.getNumDims(), newMap.getNumSymbols());
2412 }
2413 }
2414
2415 // Internal constraint variables (dimOp, dimOpBound, resultDimStart, etc.)
2416 // have no associated SSA values (null Value()). Replace their corresponding
2417 // dim/symbol positions in newMap with constant 0 and compact newOperands.
2418 // These positions should be unreferenced in newMap (the bound was computed
2419 // in terms of the original operands only), so replacing with 0 is safe.
2420 // This prevents canonicalizeMapAndOperands from using null Values as
2421 // DenseMap keys, which would cause undefined behavior.
2422 {
2423 unsigned numDims = newMap.getNumDims();
2424 unsigned numSyms = newMap.getNumSymbols();
2425 SmallVector<AffineExpr> dimReplacements(numDims), symReplacements(numSyms);
2426 SmallVector<Value> filteredOperands;
2427 filteredOperands.reserve(newOperands.size());
2428 unsigned newDim = 0;
2429 for (unsigned i = 0; i < numDims; ++i) {
2430 if (newOperands[i]) {
2431 dimReplacements[i] = getAffineDimExpr(newDim++, ctx);
2432 filteredOperands.push_back(newOperands[i]);
2433 } else {
2434 assert(!newMap.isFunctionOfDim(i) &&
2435 "null-valued dim operand referenced in bound map");
2436 dimReplacements[i] = getAffineConstantExpr(0, ctx);
2437 }
2438 }
2439 unsigned newSym = 0;
2440 for (unsigned i = 0; i < numSyms; ++i) {
2441 if (newOperands[numDims + i]) {
2442 symReplacements[i] = getAffineSymbolExpr(newSym++, ctx);
2443 filteredOperands.push_back(newOperands[numDims + i]);
2444 } else {
2445 assert(!newMap.isFunctionOfSymbol(i) &&
2446 "null-valued symbol operand referenced in bound map");
2447 symReplacements[i] = getAffineConstantExpr(0, ctx);
2448 }
2449 }
2450 newMap = newMap.replaceDimsAndSymbols(dimReplacements, symReplacements,
2451 newDim, newSym);
2452 newOperands = std::move(filteredOperands);
2453 }
2454
2455 affine::canonicalizeMapAndOperands(&newMap, &newOperands);
2456 return AffineValueMap(newMap, newOperands);
2457}
2458
2460 Operation *b) {
2461 Region *aScope = getAffineAnalysisScope(a);
2462 Region *bScope = getAffineAnalysisScope(b);
2463 if (aScope != bScope)
2464 return nullptr;
2465
2466 // Get the block ancestry of `op` while stopping at the affine scope `aScope`
2467 // and store them in `ancestry`.
2468 auto getBlockAncestry = [&](Operation *op,
2469 SmallVectorImpl<Block *> &ancestry) {
2470 Operation *curOp = op;
2471 do {
2472 ancestry.push_back(curOp->getBlock());
2473 if (curOp->getParentRegion() == aScope)
2474 break;
2475 curOp = curOp->getParentOp();
2476 } while (curOp);
2477 assert(curOp && "can't reach root op without passing through affine scope");
2478 std::reverse(ancestry.begin(), ancestry.end());
2479 };
2480
2481 SmallVector<Block *, 4> aAncestors, bAncestors;
2482 getBlockAncestry(a, aAncestors);
2483 getBlockAncestry(b, bAncestors);
2484 assert(!aAncestors.empty() && !bAncestors.empty() &&
2485 "at least one Block ancestor expected");
2486
2487 Block *innermostCommonBlock = nullptr;
2488 for (unsigned a = 0, b = 0, e = aAncestors.size(), f = bAncestors.size();
2489 a < e && b < f; ++a, ++b) {
2490 if (aAncestors[a] != bAncestors[b])
2491 break;
2492 innermostCommonBlock = aAncestors[a];
2493 }
2494 return innermostCommonBlock;
2495}
return success()
static std::optional< uint64_t > getConstDifference(AffineMap lbMap, AffineMap ubMap)
Returns the number of iterations the slice bounded below by lbMap and above by ubMap runs for,...
Definition Utils.cpp:1813
static bool mayAccessMemRef(Operation *op, Value memref)
Returns true if op may read from or write to memref.
Definition Utils.cpp:247
static void findInstPosition(Operation *op, Block *limitBlock, SmallVectorImpl< unsigned > *positions)
Definition Utils.cpp:1550
static Node * addNodeToMDG(Operation *nodeOp, MemRefDependenceGraph &mdg, DenseMap< Value, SetVector< unsigned > > &memrefAccesses)
Add op to MDG creating a new node and adding its memory accesses (affine or non-affine to memrefAcces...
Definition Utils.cpp:192
static bool mayDependence(const Node &srcNode, const Node &dstNode, Value memref)
Returns true if there may be a dependence on memref from srcNode's memory ops to dstNode's memory ops...
Definition Utils.cpp:276
const char *const kSliceFusionBarrierAttrName
Definition Utils.cpp:1887
static LogicalResult addMissingLoopIVBounds(SmallPtrSet< Value, 8 > &ivs, FlatAffineValueConstraints *cst)
Definition Utils.cpp:1592
static void getEffectedValues(Operation *op, SmallVectorImpl< Value > &values)
Returns the values that this op has a memref effect of type EffectTys on, not considering recursive e...
Definition Utils.cpp:166
MemRefDependenceGraph::Node Node
Definition Utils.cpp:39
static void unpackOptionalValues(ArrayRef< std::optional< Value > > source, SmallVector< Value > &target)
Definition Utils.cpp:2240
static AffineMap addConstToResults(AffineMap map, int64_t val)
Add val to each result of map.
Definition Utils.cpp:2268
static LogicalResult alignAndAddBound(FlatAffineValueConstraints &constraints, BoundType type, unsigned pos, AffineMap map, ValueRange operands)
Bound an identifier pos in a given FlatAffineValueConstraints with constraints drawn from an affine m...
Definition Utils.cpp:2253
static Operation * getInstAtPosition(ArrayRef< unsigned > positions, unsigned level, Block *block)
Definition Utils.cpp:1567
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be inserted(the insertion happens right before the *insertion point). Since `begin` can itself be invalidated due to the memref *rewriting done from this method
template bool mlir::hasEffect< MemoryEffects::Free >(Operation *)
A dimensional identifier appearing in an affine expression.
Definition AffineExpr.h:223
Base type for affine expression.
Definition AffineExpr.h:68
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
MLIRContext * getContext() const
bool isFunctionOfDim(unsigned position) const
Return true if any affine expression involves AffineDimExpr position.
Definition AffineMap.h:221
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
AffineMap shiftDims(unsigned shift, unsigned offset=0) const
Replace dims[offset ... numDims) by dims[offset + shift ... shift + numDims).
Definition AffineMap.h:267
unsigned getNumSymbols() const
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
bool isFunctionOfSymbol(unsigned position) const
Return true if any affine expression involves AffineSymbolExpr position.
Definition AffineMap.h:228
unsigned getNumResults() const
AffineMap replaceDimsAndSymbols(ArrayRef< AffineExpr > dimReplacements, ArrayRef< AffineExpr > symReplacements, unsigned numResultDims, unsigned numResultSyms) const
This method substitutes any uses of dimensions and symbols (e.g.
unsigned getNumInputs() const
AffineExpr getResult(unsigned idx) const
AffineMap replace(AffineExpr expr, AffineExpr replacement, unsigned numResultDims, unsigned numResultSyms) const
Sparse replace method.
AffineMap getSubMap(ArrayRef< unsigned > resultPos) const
Returns the map consisting of the resultPos subset.
Block represents an ordered list of Operations.
Definition Block.h:34
OpListType::iterator iterator
Definition Block.h:165
RetT walk(FnT &&callback)
Walk all nested operations, blocks (including this block) or regions, depending on the type of callba...
Definition Block.h:318
iterator begin()
Definition Block.h:168
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
Definition Block.cpp:31
This class is a general helper class for creating context-global objects like types,...
Definition Builders.h:51
AffineExpr getAffineSymbolExpr(unsigned position)
Definition Builders.cpp:377
AffineExpr getAffineConstantExpr(int64_t constant)
Definition Builders.cpp:381
AffineExpr getAffineDimExpr(unsigned position)
Definition Builders.cpp:373
void getSliceBounds(unsigned offset, unsigned num, MLIRContext *context, SmallVectorImpl< AffineMap > *lbMaps, SmallVectorImpl< AffineMap > *ubMaps, bool closedUB=false, bool allowMultiResultUB=false)
Computes the lower and upper bounds of the first num dimensional variables (starting at offset) as an...
std::optional< int64_t > getConstantBoundOnDimSize(MLIRContext *context, unsigned pos, AffineMap *lb=nullptr, AffineMap *ub=nullptr, unsigned *minLbPos=nullptr, unsigned *minUbPos=nullptr) const
Returns a non-negative constant bound on the extent (upper bound - lower bound) of the specified vari...
FlatLinearValueConstraints represents an extension of FlatLinearConstraints where each non-local vari...
LogicalResult unionBoundingBox(const FlatLinearValueConstraints &other)
Updates the constraints to be the smallest bounding (enclosing) box that contains the points of this ...
void mergeAndAlignVarsWithOther(unsigned offset, FlatLinearValueConstraints *other)
Merge and align the variables of this and other starting at offset, so that both constraint systems g...
Value getValue(unsigned pos) const
Returns the Value associated with the pos^th variable.
void projectOut(Value val)
Projects out the variable that is associate with Value.
bool containsVar(Value val) const
Returns true if a variable with the specified Value exists, false otherwise.
void addBound(presburger::BoundType type, Value val, int64_t value)
Adds a constant bound for the variable associated with the given Value.
bool areVarsAlignedWithOther(const FlatLinearConstraints &other)
Returns true if this constraint system and other are in the same space, i.e., if they are associated ...
void getValues(unsigned start, unsigned end, SmallVectorImpl< Value > *values) const
Returns the Values associated with variables in range [start, end).
SmallVector< std::optional< Value > > getMaybeValues() const
An integer set representing a conjunction of one or more affine equalities and inequalities.
Definition IntegerSet.h:44
unsigned getNumDims() const
MLIRContext * getContext() const
static IntegerSet getEmptySet(unsigned numDims, unsigned numSymbols, MLIRContext *context)
Definition IntegerSet.h:56
unsigned getNumSymbols() const
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class helps build Operations.
Definition Builders.h:210
A trait of region holding operations that defines a new scope for polyhedral optimization purposes.
This trait indicates that the memory effects of an operation includes the effects of operations neste...
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:801
bool isBeforeInBlock(Operation *other)
Given an operation 'other' that is within the same parent block, return whether the current operation...
InFlightDiagnostic emitWarning(const Twine &message={})
Emit a warning about this operation, reporting up to any diagnostic handlers that may be listening.
Block * getBlock()
Returns the operation block that contains this operation.
Definition Operation.h:230
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
Definition Operation.h:729
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
Definition Operation.h:849
user_range getUsers()
Returns a range of all users.
Definition Operation.h:925
result_range getResults()
Definition Operation.h:440
Region * getParentRegion()
Returns the region to which the instruction belongs.
Definition Operation.h:247
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
A utility result that is used to signal how to proceed with an ongoing walk:
Definition WalkResult.h:29
static WalkResult advance()
Definition WalkResult.h:47
An AffineValueMap is an affine map plus its ML value operands and results for analysis purposes.
Value getOperand(unsigned i) const
unsigned getNumOperands() const
FlatAffineValueConstraints is an extension of FlatLinearValueConstraints with helper functions for Af...
LogicalResult addBound(presburger::BoundType type, unsigned pos, AffineMap boundMap, ValueRange operands)
Adds a bound for the variable at the specified position with constraints being drawn from the specifi...
void convertLoopIVSymbolsToDims()
Changes all symbol variables which are loop IVs to dim variables.
LogicalResult addDomainFromSliceMaps(ArrayRef< AffineMap > lbMaps, ArrayRef< AffineMap > ubMaps, ArrayRef< Value > operands)
Adds constraints (lower and upper bounds) for each loop in the loop nest described by the bound maps ...
LogicalResult addAffineForOpDomain(AffineForOp forOp)
Adds constraints (lower and upper bounds) for the specified 'affine.for' operation's Value using IR i...
static FailureOr< FlatAffineValueConstraints > create(IntegerSet set, ValueRange operands={})
Creates an affine constraint system from an IntegerSet.
LogicalResult addSliceBounds(ArrayRef< Value > values, ArrayRef< AffineMap > lbMaps, ArrayRef< AffineMap > ubMaps, ArrayRef< Value > operands)
Adds slice lower bounds represented by lower bounds in lbMaps and upper bounds in ubMaps to each vari...
bool isEmpty() const
Checks for emptiness by performing variable elimination on all variables, running the GCD test on eac...
unsigned getNumCols() const
Returns the number of columns in the constraint system.
void addInequality(ArrayRef< DynamicAPInt > inEq)
Adds an inequality (>= 0) from the coefficients specified in inEq.
std::optional< int64_t > getConstantBound64(BoundType type, unsigned pos) const
The same, but casts to int64_t.
bool isIntegerEmpty() const
Return true if all the sets in the union are known to be integer empty false otherwise.
PresburgerSet subtract(const PresburgerRelation &set) const
IntegerSet simplifyIntegerSet(IntegerSet set)
Simplify the integer set by simplifying the underlying affine expressions by flattening and some simp...
Definition Utils.cpp:2224
void getEnclosingAffineOps(Operation &op, SmallVectorImpl< Operation * > *ops)
Populates 'ops' with affine operations enclosing op ordered from outermost to innermost while stoppin...
Definition Utils.cpp:876
SliceComputationResult computeSliceUnion(ArrayRef< Operation * > opsA, ArrayRef< Operation * > opsB, unsigned loopDepth, unsigned numCommonLoops, bool isBackwardSlice, ComputationSliceState *sliceUnion)
Computes in 'sliceUnion' the union of all slice bounds computed at 'loopDepth' between all dependent ...
Definition Utils.cpp:1639
bool isAffineInductionVar(Value val)
Returns true if the provided value is the induction variable of an AffineForOp or AffineParallelOp.
bool isLoopParallelAndContainsReduction(AffineForOp forOp)
Returns whether a loop is a parallel loop and contains a reduction loop.
Definition Utils.cpp:2206
unsigned getNumCommonSurroundingLoops(Operation &a, Operation &b)
Returns the number of surrounding loops common to both A and B.
Definition Utils.cpp:2132
AffineForOp getForInductionVarOwner(Value val)
Returns the loop parent of an induction variable.
void getAffineIVs(Operation &op, SmallVectorImpl< Value > &ivs)
Populates 'ivs' with IVs of the surrounding affine.for and affine.parallel ops ordered from the outer...
Definition Utils.cpp:2115
void getSequentialLoops(AffineForOp forOp, llvm::SmallDenseSet< Value, 8 > *sequentialLoops)
Returns in 'sequentialLoops' all sequential loops in loop nest rooted at 'forOp'.
Definition Utils.cpp:2215
DependenceResult checkMemrefAccessDependence(const MemRefAccess &srcAccess, const MemRefAccess &dstAccess, unsigned loopDepth, FlatAffineValueConstraints *dependenceConstraints=nullptr, SmallVector< DependenceComponent, 2 > *dependenceComponents=nullptr, bool allowRAR=false)
void canonicalizeMapAndOperands(AffineMap *map, SmallVectorImpl< Value > *operands)
Modifies both map and operands in-place so as to:
bool isAffineForInductionVar(Value val)
Returns true if the provided value is the induction variable of an AffineForOp.
Region * getAffineAnalysisScope(Operation *op)
Returns the closest region enclosing op that is held by a non-affine operation; nullptr if there is n...
void getAffineForIVs(Operation &op, SmallVectorImpl< AffineForOp > *loops)
Populates 'loops' with IVs of the affine.for ops surrounding 'op' ordered from the outermost 'affine....
Definition Utils.cpp:862
std::optional< int64_t > getMemoryFootprintBytes(AffineForOp forOp, int memorySpace=-1)
Gets the memory footprint of all data touched in the specified memory space in bytes; if the memory s...
Definition Utils.cpp:2197
unsigned getInnermostCommonLoopDepth(ArrayRef< Operation * > ops, SmallVectorImpl< AffineForOp > *surroundingLoops=nullptr)
Returns the innermost common loop depth for the set of operations in 'ops'.
Definition Utils.cpp:1608
bool isValidSymbol(Value value)
Returns true if the given value can be used as a symbol in the region of the closest surrounding op t...
void getComputationSliceState(Operation *depSourceOp, Operation *depSinkOp, const FlatAffineValueConstraints &dependenceConstraints, unsigned loopDepth, bool isBackwardSlice, ComputationSliceState *sliceState)
Computes the computation slice loop bounds for one loop nest as affine maps of the other loop nest's ...
Definition Utils.cpp:1892
AffineParallelOp getAffineParallelInductionVarOwner(Value val)
Returns true if the provided value is among the induction variables of an AffineParallelOp.
std::optional< uint64_t > getIntOrFloatMemRefSizeInBytes(MemRefType memRefType)
Returns the size of a memref with element type int or float in bytes if it's statically shaped,...
Definition Utils.cpp:1470
unsigned getNestingDepth(Operation *op)
Returns the nesting depth of this operation, i.e., the number of loops surrounding this operation.
Definition Utils.cpp:2087
uint64_t getSliceIterationCount(const llvm::SmallDenseMap< Operation *, uint64_t, 8 > &sliceTripCountMap)
Return the number of iterations for the slicetripCountMap provided.
Definition Utils.cpp:1878
LogicalResult boundCheckLoadOrStoreOp(LoadOrStoreOpPointer loadOrStoreOp, bool emitError=true)
Checks a load or store op for an out of bound access; returns failure if the access is out of bounds ...
mlir::Block * findInnermostCommonBlockInScope(mlir::Operation *a, mlir::Operation *b)
Find the innermost common Block of a and b in the affine scope that a and b are part of.
Definition Utils.cpp:2459
bool noDependence(DependenceResult result)
Returns true if the provided DependenceResult corresponds to the absence of a dependence.
bool buildSliceTripCountMap(const ComputationSliceState &slice, llvm::SmallDenseMap< Operation *, uint64_t, 8 > *tripCountMap)
Builds a map 'tripCountMap' from AffineForOp to constant trip count for loop nest surrounding represe...
Definition Utils.cpp:1840
AffineForOp insertBackwardComputationSlice(Operation *srcOpInst, Operation *dstOpInst, unsigned dstLoopDepth, ComputationSliceState *sliceState)
Creates a clone of the computation contained in the loop nest surrounding 'srcOpInst',...
Definition Utils.cpp:2006
FailureOr< AffineValueMap > simplifyConstrainedMinMaxOp(Operation *op, FlatAffineValueConstraints constraints)
Try to simplify the given affine.min or affine.max op to an affine map with a single result and opera...
Definition Utils.cpp:2313
std::optional< int64_t > getMemRefIntOrFloatEltSizeInBytes(MemRefType memRefType)
Returns the memref's element type's size in bytes where the elemental type is an int or float or a ve...
Definition Utils.cpp:1426
BoundType
The type of bound: equal, lower bound or upper bound.
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
llvm::DenseSet< ValueT, ValueInfoT > DenseSet
Definition LLVM.h:122
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
AffineMap alignAffineMapWithValues(AffineMap map, ValueRange operands, ValueRange dims, ValueRange syms, SmallVector< Value > *newSyms=nullptr)
Re-indexes the dimensions and symbols of an affine map with given operands values to align with dims ...
llvm::SetVector< T, Vector, Set, N > SetVector
Definition LLVM.h:125
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
bool hasEffect(Operation *op)
Returns "true" if op has an effect of type EffectTy.
AffineExpr simplifyAffineExpr(AffineExpr expr, unsigned numDims, unsigned numSymbols)
Simplify an affine expression by flattening and some amount of simple analysis.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
AffineExpr getAffineSymbolExpr(unsigned position, MLIRContext *context)
ComputationSliceState aggregates loop IVs, loop bound AffineMaps and their associated operands for a ...
Definition Utils.h:318
std::optional< bool > isSliceValid() const
Checks the validity of the slice computed.
Definition Utils.cpp:1047
SmallVector< Value, 4 > ivs
Definition Utils.h:321
LogicalResult getAsConstraints(FlatAffineValueConstraints *cst) const
Definition Utils.cpp:909
LogicalResult getSourceAsConstraints(FlatAffineValueConstraints &cst) const
Adds to 'cst' constraints which represent the original loop bounds on 'ivs' in 'this'.
Definition Utils.cpp:893
std::vector< SmallVector< Value, 4 > > ubOperands
Definition Utils.h:329
SmallVector< AffineMap, 4 > ubs
Definition Utils.h:325
std::optional< bool > isMaximal() const
Returns true if the computation slice encloses all the iterations of the sliced loop nest.
Definition Utils.cpp:1113
SmallVector< AffineMap, 4 > lbs
Definition Utils.h:323
std::vector< SmallVector< Value, 4 > > lbOperands
Definition Utils.h:327
Checks whether two accesses to the same memref access the same element.
SmallVector< Operation *, 4 > memrefFrees
Definition Utils.h:49
SmallVector< AffineForOp, 4 > forOps
Definition Utils.h:39
SmallVector< Operation *, 4 > loadOpInsts
Definition Utils.h:41
SmallVector< Operation *, 4 > memrefStores
Definition Utils.h:47
void collect(Operation *opToWalk)
Definition Utils.cpp:44
SmallVector< Operation *, 4 > memrefLoads
Definition Utils.h:45
SmallVector< Operation *, 4 > storeOpInsts
Definition Utils.h:43
Encapsulates a memref load or store access information.
unsigned getRank() const
Definition Utils.cpp:2077
SmallVector< Value, 4 > indices
void getAccessMap(AffineValueMap *accessMap) const
Populates 'accessMap' with composition of AffineApplyOps reachable from 'indices'.
MemRefAccess(Operation *memOp)
Constructs a MemRefAccess from an affine read/write operation.
Definition Utils.cpp:2062
bool operator==(const MemRefAccess &rhs) const
Equal if both affine accesses can be proved to be equivalent at compile time (considering the memrefs...
Definition Utils.cpp:2105
void getStoreOpsForMemref(Value memref, SmallVectorImpl< Operation * > *storeOps) const
Definition Utils.cpp:131
SmallVector< Operation *, 4 > loads
Definition Utils.h:73
SmallVector< Operation *, 4 > stores
Definition Utils.h:77
void getLoadAndStoreMemrefSet(DenseSet< Value > *loadAndStoreMemrefSet) const
Definition Utils.cpp:150
unsigned hasFree(Value memref) const
Definition Utils.cpp:124
SmallVector< Operation *, 4 > memrefLoads
Definition Utils.h:75
SmallVector< Operation *, 4 > memrefStores
Definition Utils.h:79
unsigned getLoadOpCount(Value memref) const
Definition Utils.cpp:79
unsigned getStoreOpCount(Value memref) const
Definition Utils.cpp:94
unsigned hasStore(Value memref) const
Returns true if there exists an operation with a write memory effect to memref in this node.
Definition Utils.cpp:110
void getLoadOpsForMemref(Value memref, SmallVectorImpl< Operation * > *loadOps) const
Definition Utils.cpp:140
SmallVector< Operation *, 4 > memrefFrees
Definition Utils.h:81
DenseMap< unsigned, SmallVector< Edge, 2 > > outEdges
Definition Utils.h:147
Block & block
The block for which this graph is created to perform fusion.
Definition Utils.h:270
unsigned addNode(Operation *op)
Definition Utils.cpp:478
bool writesToLiveInOrEscapingMemrefs(unsigned id) const
Definition Utils.cpp:508
void removeEdge(unsigned srcId, unsigned dstId, Value value)
Definition Utils.cpp:554
void addEdge(unsigned srcId, unsigned dstId, Value value)
Definition Utils.cpp:543
DenseMap< unsigned, Node > nodes
Definition Utils.h:141
void gatherDefiningNodes(unsigned id, DenseSet< unsigned > &definingNodes) const
Return all nodes which define SSA values used in node 'id'.
Definition Utils.cpp:642
bool hasDependencePath(unsigned srcId, unsigned dstId) const
Definition Utils.cpp:581
void clearNodeLoadAndStores(unsigned id)
Definition Utils.cpp:808
const Node * getForOpNode(AffineForOp forOp) const
Definition Utils.cpp:470
Operation * getFusedLoopNestInsertionPoint(unsigned srcId, unsigned dstId) const
Definition Utils.cpp:656
void updateEdges(unsigned srcId, unsigned dstId, const DenseSet< Value > &privateMemRefs, bool removeSrcId)
Definition Utils.cpp:731
DenseMap< unsigned, SmallVector< Edge, 2 > > inEdges
Definition Utils.h:144
void forEachMemRefInputEdge(unsigned id, const std::function< void(Edge)> &callback)
Definition Utils.cpp:816
bool init(bool fullAffineDependences=true)
Definition Utils.cpp:347
unsigned getOutEdgeCount(unsigned id, Value memref=nullptr) const
Definition Utils.cpp:632
const Node * getNode(unsigned id) const
Definition Utils.cpp:463
void forEachMemRefOutputEdge(unsigned id, const std::function< void(Edge)> &callback)
Definition Utils.cpp:824
void forEachMemRefEdge(ArrayRef< Edge > edges, const std::function< void(Edge)> &callback)
Definition Utils.cpp:832
void addToNode(unsigned id, ArrayRef< Operation * > loads, ArrayRef< Operation * > stores, ArrayRef< Operation * > memrefLoads, ArrayRef< Operation * > memrefStores, ArrayRef< Operation * > memrefFrees)
Definition Utils.cpp:795
unsigned getIncomingMemRefAccesses(unsigned id, Value memref) const
Definition Utils.cpp:616
bool hasEdge(unsigned srcId, unsigned dstId, Value value=nullptr) const
Definition Utils.cpp:528
void print(raw_ostream &os) const
Definition Utils.cpp:844
DenseMap< Value, unsigned > memrefEdgeCount
Definition Utils.h:150
A region of a memref's data space; this is typically constructed by analyzing load/store op's on this...
Definition Utils.h:489
std::optional< int64_t > getConstantBoundingSizeAndShape(SmallVectorImpl< int64_t > *shape=nullptr, SmallVectorImpl< AffineMap > *lbs=nullptr) const
Returns a constant upper bound on the number of elements in this region if bounded by a known constan...
Definition Utils.cpp:1167
unsigned getRank() const
Returns the rank of the memref that this region corresponds to.
Definition Utils.cpp:1163
FlatAffineValueConstraints cst
Region (data space) of the memref accessed.
Definition Utils.h:586
LogicalResult compute(Operation *op, unsigned loopDepth, const ComputationSliceState *sliceState=nullptr, bool addMemRefDimBounds=true, bool dropLocalVars=true, bool dropOuterIVs=true)
Computes the memory region accessed by this memref with the region represented as constraints symboli...
Definition Utils.cpp:1264
void getLowerAndUpperBound(unsigned pos, AffineMap &lbMap, AffineMap &ubMap) const
Gets the lower and upper bound map for the dimensional variable at pos.
Definition Utils.cpp:1223
std::optional< int64_t > getRegionSize()
Returns the size of this MemRefRegion in bytes.
Definition Utils.cpp:1445
LogicalResult unionBoundingBox(const MemRefRegion &other)
Definition Utils.cpp:1242
bool write
Read or write.
Definition Utils.h:573
FlatAffineValueConstraints * getConstraints()
Definition Utils.h:535
Value memref
Memref that this region corresponds to.
Definition Utils.h:570
MemRefRegion(Location loc)
Definition Utils.h:490
Enumerates different result statuses of slice computation by computeSliceUnion
Definition Utils.h:305
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.