54 static std::unique_ptr<Context>
57 return std::make_unique<Context>();
60 std::unique_ptr<SCEVChainRecurrence>
67 return SCEV::CloneAs<SCEVChainRecurrence>(*it->second);
73 const auto it =
SCEVMap_.find(&output);
74 if (it ==
SCEVMap_.end() || !it->second)
77 return it->second->Clone();
83 ChrecMap_.insert_or_assign(&output, SCEV::CloneAs<SCEVChainRecurrence>(*chrec));
86 const std::unordered_map<rvsdg::Output *, std::unique_ptr<SCEVChainRecurrence>> &
92 const std::unordered_map<rvsdg::Output *, std::unique_ptr<SCEV>> &
108 if (chrec->GetOperands().size() == n + 1 && !
IsUnknown(*chrec))
135 SCEVMap_.insert_or_assign(&output, scev->Clone());
162 const std::unordered_map<const rvsdg::ThetaNode *, size_t> &
169 std::unordered_map<rvsdg::Output *, std::unique_ptr<SCEVChainRecurrence>>
ChrecMap_;
170 std::unordered_map<rvsdg::Output *, std::unique_ptr<SCEV>>
SCEVMap_;
197 AddMeasurement(Label::NumTotalInductionVariables, context.GetNumTotalInductionVariables());
199 Label::NumConstantInductionVariables,
200 context.GetNumInductionVariablesWithOrder(0));
202 Label::NumFirstOrderInductionVariables,
203 context.GetNumInductionVariablesWithOrder(1));
205 Label::NumSecondOrderInductionVariables,
206 context.GetNumInductionVariablesWithOrder(2));
207 AddMeasurement(Label::NumLoopVariablesTotal, context.GetNumTotalLoopVars());
217 for (
auto & [thetaNode, tripCount] : tripCountMap)
223 s +=
"ID(" + std::to_string(thetaNode->subregion()->getRegionId())
224 +
")=" + std::to_string(tripCount);
229 static std::unique_ptr<Statistics>
232 return std::make_unique<Statistics>(sourceFile);
237 : rvsdg::Transformation(
"ScalarEvolution")
245 std::unordered_map<const rvsdg::Output *, std::unique_ptr<SCEVChainRecurrence>> mapCopy{};
246 for (
auto & [output, chrec] : Context_->GetChrecMap())
248 mapCopy.emplace(output, SCEV::CloneAs<SCEVChainRecurrence>(*chrec));
253std::unordered_map<const rvsdg::Output *, std::unique_ptr<SCEV>>
256 std::unordered_map<const rvsdg::Output *, std::unique_ptr<SCEV>> mapCopy{};
257 for (
auto & [output, scev] :
Context_->GetSCEVMap())
259 mapCopy.emplace(output, scev->Clone());
264std::unordered_map<const rvsdg::ThetaNode *, size_t>
290 for (
auto & node : region.
Nodes())
294 for (
auto & subregion : structuralNode->Subregions())
302 for (
const auto loopVar : thetaNode->GetLoopVars())
315 if (tripCount.has_value())
316 Context_->SetTripCount(*thetaNode, *tripCount);
325 if (
auto constantStep =
dynamic_cast<const SCEVConstant *
>(&stepSCEV))
327 return constantStep->GetValue() < 0;
333 const auto start =
dynamic_cast<const SCEVConstant *
>(recurrenceStep->GetStartValue());
334 auto stepPtr = recurrenceStep->GetStep();
335 const auto step =
dynamic_cast<const SCEVConstant *
>(stepPtr->get());
338 throw std::logic_error(
"Step can only contain constant SCEVs!");
340 const auto a = start->GetValue();
341 const auto b = step->GetValue();
343 return a <= 0 && b <= 0 && !(a == 0 && b == 0);
345 throw std::logic_error(
"Wrong type for step!");
351 if (
auto constantStep =
dynamic_cast<const SCEVConstant *
>(&stepSCEV))
353 return constantStep->GetValue() > 0;
359 const auto start =
dynamic_cast<const SCEVConstant *
>(recurrenceStep->GetStartValue());
360 auto stepPtr = recurrenceStep->GetStep();
361 const auto step =
dynamic_cast<const SCEVConstant *
>(stepPtr->get());
364 throw std::logic_error(
"Step can only contain constant SCEVs!");
366 const auto a = start->GetValue();
367 const auto b = step->GetValue();
369 return a >= 0 && b >= 0 && !(a == 0 && b == 0);
371 throw std::logic_error(
"Wrong type for step!");
377 if (
auto constantStep =
dynamic_cast<const SCEVConstant *
>(&stepSCEV))
379 return constantStep->GetValue() == 0;
385 const auto start =
dynamic_cast<const SCEVConstant *
>(recurrenceStep->GetStartValue());
386 auto stepPtr = recurrenceStep->GetStep();
387 const auto step =
dynamic_cast<const SCEVConstant *
>(stepPtr->get());
390 throw std::logic_error(
"Step can only contain constant SCEVs!");
392 const auto a = start->GetValue();
393 const auto b = step->GetValue();
395 return a == 0 && b == 0;
397 throw std::logic_error(
"Wrong type for step!");
404 const auto & [node, matchOperation] =
411 const auto origin = node->input(0)->origin();
416 const auto * comparisonOperation = &comparisonNode->GetOperation();
429 auto & lhs = *comparisonNode->input(0)->origin();
430 auto & rhs = *comparisonNode->input(1)->origin();
431 auto lhsChrec =
Context_->TryGetChrecForOutput(lhs);
432 auto rhsChrec =
Context_->TryGetChrecForOutput(rhs);
441 std::unique_ptr<SCEVChainRecurrence> chrec{};
445 const auto constantSCEV =
dynamic_cast<SCEVConstant *
>(lhsChrec->GetOperand(0));
450 chrec = SCEV::CloneAs<SCEVChainRecurrence>(*rhsChrec);
454 const auto constantSCEV =
dynamic_cast<SCEVConstant *
>(rhsChrec->GetOperand(0));
459 chrec = SCEV::CloneAs<SCEVChainRecurrence>(*lhsChrec);
474 for (
const auto op : chrec->GetOperands())
484 const auto start =
dynamic_cast<const SCEVConstant *
>(chrec->GetStartValue())->GetValue();
485 const auto stepOpt = chrec->GetStep();
489 const auto & stepSCEV = **stepOpt;
499 const auto backedgeTakenCount =
501 if (backedgeTakenCount.has_value())
504 return *backedgeTakenCount + 1;
515 const auto backedgeTakenCount =
517 if (backedgeTakenCount.has_value())
519 return *backedgeTakenCount + 1;
530 const auto backedgeTakenCount =
532 if (backedgeTakenCount.has_value())
534 return *backedgeTakenCount + 1;
545 const auto backedgeTakenCount =
547 if (backedgeTakenCount.has_value())
549 return *backedgeTakenCount + 1;
560 const auto step =
dynamic_cast<const SCEVConstant *
>(&stepSCEV)->GetValue();
563 const auto backedgeTakenCount =
566 if (start <= bound && (bound - start) % step == 0)
567 return *backedgeTakenCount + 1;
571 const auto backedgeTakenCount =
573 if (start >= bound && (bound - start) % step == 0)
574 return *backedgeTakenCount + 1;
613 const auto stepOpt = chrec.
GetStep();
617 const auto & stepSCEV = *stepOpt;
628 const auto stepConstant =
dynamic_cast<const SCEVConstant *
>(stepSCEV.get());
629 const auto step = stepConstant->
GetValue();
633 size_t result = std::ceil(
static_cast<double>(bound - start) / step);
635 if (isEqualsComparison)
639 if ((bound - start) % step == 0)
652 const int64_t stepFirst =
653 dynamic_cast<const SCEVConstant *
>(stepRecurrence->GetStartValue())->GetValue();
655 const int64_t stepSecond =
656 dynamic_cast<const SCEVConstant *
>(stepRecurrence->GetStep()->get())->GetValue();
671 const int64_t a = stepSecond;
672 const int64_t b = 2 * stepFirst - stepSecond;
673 const int64_t c = 2 * (start - bound);
676 if (!quadraticResult.has_value())
679 size_t result = *quadraticResult;
681 if (isEqualsComparison)
685 const int64_t valueAtResult =
686 start + result * stepFirst + result * (result - 1) / 2 * stepSecond;
687 if (valueAtResult == bound)
706 const auto d = b * b - 4 * a * c;
712 int64_t sq = std::floor(std::sqrt(d));
715 const bool inexactSq = (sq * sq != d);
733 x = (-b - (sq + (inexactSq ? 1 : 0))) / (2 * a);
734 rem = (-b - sq) % (2 * a);
738 x = (-b + sq) / (2 * a);
739 rem = (-b + sq) % (2 * a);
747 if (!inexactSq && rem == 0)
754 const int64_t valueAtX = (a * x + b) * x + c;
755 const int64_t valueAtXPlusOne = (a * (x + 1) + b) * (x + 1) + c;
757 const bool signChange =
758 ((valueAtX < 0) != (valueAtXPlusOne < 0)) || ((valueAtX == 0) != (valueAtXPlusOne == 0));
775 std::vector<std::pair<rvsdg::Output *, std::unique_ptr<SCEV>>> pending;
776 for (
auto & [output, chrec] :
Context_->GetChrecMap())
780 pending.emplace_back(output, std::move(*newSCEV));
785 for (
auto & [output, scev] : pending)
790 Context_->InsertChrec(*output, SCEV::CloneAs<SCEVChainRecurrence>(*chrec));
796 Context_->InsertSCEV(*output, std::move(scev));
802std::optional<std::unique_ptr<SCEV>>
805 if (
const auto initSCEV =
dynamic_cast<const SCEVInit *
>(&scev))
809 const auto & initPrePointer = initSCEV->GetPrePointer();
812 const auto correspondingInput = innerTheta->MapPreLoopVar(initPrePointer).input;
814 if (
const auto originSCEV =
Context_->TryGetSCEVForOutput(inputOrigin))
819 const auto outerTheta =
820 thetaParent ? thetaParent
821 : util::assertedCast<rvsdg::ThetaNode>(inputOrigin.region()->node());
826 return chrec->Clone();
830 if (
const auto nArySCEV =
dynamic_cast<const SCEVNAryExpr *
>(&scev))
834 auto clone = SCEV::CloneAs<SCEVNAryExpr>(*nArySCEV);
835 const auto operands = nArySCEV->GetOperands();
836 bool changed =
false;
837 for (
size_t i = 0; i < operands.size(); ++i)
845 clone->ReplaceOperand(i, std::move(*result));
879 const auto post = loopVar.post;
883 Context_->InsertSCEV(*loopVar.output, scev);
889 for (
const auto & [output, deps] : dependencyGraph)
892 validOutputs.
insert(output);
897 auto filteredDependencyGraph = dependencyGraph;
898 for (
auto it = filteredDependencyGraph.begin(); it != filteredDependencyGraph.end();)
900 if (!validOutputs.Contains(it->first))
902 for (
auto & [node, deps] : filteredDependencyGraph)
903 deps.erase(it->first);
904 it = filteredDependencyGraph.erase(it);
912 for (
auto output : order)
914 std::unique_ptr<SCEV> scev{};
921 scev =
Context_->TryGetSCEVForOutput(newOutput);
924 scev =
Context_->TryGetSCEVForOutput(*output);
929 Context_->InsertChrec(*output, chrec);
932 for (
auto & [output, scev] :
Context_->GetSCEVMap())
934 if (std::find(order.begin(), order.end(), output) == order.end())
936 auto unknownChainRecurrence =
938 Context_->InsertChrec(*output, unknownChainRecurrence);
946 if (
const auto existing =
Context_->TryGetSCEVForOutput(output))
947 return existing->Clone();
949 std::unique_ptr<SCEV> result{};
957 const auto & [simpleNode, simpleOperation] =
964 const auto barredInputOrigin =
977 const auto baseIndex = simpleNode->input(0)->origin();
978 JLM_ASSERT(is<PointerType>(baseIndex->Type()));
980 const auto & pointeeType = gepOp->getPointeeType();
987 std::unique_ptr<SCEV> offset =
996 const auto value = constOp->Representation().to_int();
1002 const auto lhs = simpleNode->input(0)->origin();
1003 const auto rhs = simpleNode->input(1)->origin();
1023 if (
const auto * rhsConst =
dynamic_cast<SCEVConstant *
>(rhsScev.get()))
1025 const auto shiftAmount = rhsConst->GetValue();
1034 for (
auto & input : simpleNode->Inputs())
1046 Context_->InsertSCEV(output, result);
1051std::unique_ptr<SCEV>
1054 const size_t inputIndex,
1059 if (inputIndex >= gepNode.
ninputs())
1064 const auto gepInput = gepNode.
input(inputIndex);
1065 if (
const auto arrayType =
dynamic_cast<const ArrayType *
>(&type))
1067 const auto & elementType = *arrayType->GetElementType();
1079 if (
const auto structType =
dynamic_cast<const StructType *
>(&type))
1083 if (!indexingValue.has_value())
1086 const auto & fieldType = structType->getElementType(*indexingValue);
1097 throw std::logic_error(
"Unknown GEP type!");
1106 if (
const auto placeholderSCEV =
dynamic_cast<const SCEVPlaceholder *
>(&scev))
1108 auto & dependency = placeholderSCEV->GetPrePointer();
1112 auto & depInfo = dependencies[&dependency];
1113 depInfo.operation = op;
1117 if (
const auto addSCEV =
dynamic_cast<const SCEVAddExpr *
>(&scev))
1123 if (
const auto mulSCEV =
dynamic_cast<const SCEVMulExpr *
>(&scev))
1135 for (
const auto & [output, scev] :
Context_->GetSCEVMap())
1139 theta == &thetaNode)
1143 const auto loopVar = theta->MapPreLoopVar(*output);
1144 auto newScev =
Context_->TryGetSCEVForOutput(*loopVar.post->origin());
1151 graph[output] = dependencies;
1157std::vector<rvsdg::Output *>
1160 const size_t numVertices = dependencyGraph.size();
1161 std::unordered_map<const rvsdg::Output *, int> indegree(numVertices);
1162 std::queue<rvsdg::Output *> q{};
1163 for (
auto & [node, deps] : dependencyGraph)
1165 for (
auto & dep : deps)
1167 if (
const auto ptr = dep.first; ptr == node)
1170 indegree[node] += 1;
1172 if (indegree[node] == 0)
1179 std::vector<rvsdg::Output *> result{};
1184 result.push_back(currentNode);
1186 for (
const auto & [node, deps] : dependencyGraph)
1188 if (node == currentNode)
1191 for (
const auto & dep : deps)
1193 const auto ptr = dep.first;
1196 if (ptr == currentNode)
1199 indegree[node] -= 1;
1200 if (indegree[node] == 0)
1210std::unique_ptr<SCEVChainRecurrence>
1216 if (
const auto existing =
Context_->TryGetChrecForOutput(output))
1218 return SCEV::CloneAs<SCEVChainRecurrence>(*existing);
1224 theta == &thetaNode)
1240 return stepRecurrence;
1243std::unique_ptr<SCEVChainRecurrence>
1246 const SCEV & scevTree,
1249 if (
const auto scevConstant =
dynamic_cast<const SCEVConstant *
>(&scevTree))
1254 if (
const auto scevPlaceholder =
dynamic_cast<const SCEVPlaceholder *
>(&scevTree))
1256 if (&scevPlaceholder->GetPrePointer() == &output)
1263 if (
auto storedRec =
Context_->TryGetChrecForOutput(scevPlaceholder->GetPrePointer()))
1271 if (
const auto scevAddExpr =
dynamic_cast<const SCEVAddExpr *
>(&scevTree))
1276 return SCEV::CloneAs<SCEVChainRecurrence>(
1279 if (
const auto scevMulExpr =
dynamic_cast<const SCEVMulExpr *
>(&scevTree))
1284 return SCEV::CloneAs<SCEVChainRecurrence>(
1290std::unique_ptr<SCEV>
1299 for (
size_t i = 0; i < expression.
NumOperands(); ++i)
1301 std::vector<SCEV *> ops = expression.
GetOperands();
1302 if (
dynamic_cast<const SCEVInit *
>(ops[i]))
1304 for (
size_t j = i + 1; j < expression.
NumOperands(); ++j)
1306 if (
dynamic_cast<const SCEVInit *
>(ops[j]))
1310 std::unique_ptr<SCEV> foldedOperand{};
1321 throw std::logic_error(
"Invalid n-ary SCEV expression type in FoldNAryExpression!");
1339 return expression.
Clone();
1342std::unique_ptr<SCEV>
1365 if (
const auto *lhsUnknown =
dynamic_cast<const SCEVUnknown *
>(lhsOperand),
1366 *rhsUnknown =
dynamic_cast<const SCEVUnknown *
>(rhsOperand);
1367 lhsUnknown || rhsUnknown)
1375 if (lhsChrec && rhsChrec)
1377 if (&lhsChrec->GetLoop() != &rhsChrec->GetLoop())
1380 lhsChrec->GetLoop(),
1386 const auto lhsSize = lhsChrec->NumOperands();
1387 const auto rhsSize = rhsChrec->NumOperands();
1388 for (
size_t i = 0; i < std::max(lhsSize, rhsSize); ++i)
1393 lhs = lhsChrec->GetOperand(i);
1396 rhs = rhsChrec->GetOperand(i);
1404 if (lhsChrec || rhsChrec)
1406 auto * chrec = lhsChrec ? lhsChrec : rhsChrec;
1407 auto * otherOperand = lhsChrec ? rhsOperand : lhsOperand;
1410 if (
const auto constant =
dynamic_cast<const SCEVConstant *
>(otherOperand))
1414 return chrec->
Clone();
1418 const auto chrecOperands = chrec->GetOperands();
1420 bool isFirst =
true;
1421 for (
const auto operand : chrecOperands)
1431 newChrec->AddOperand(operand->Clone());
1437 const auto lhsNAryMulExpr =
dynamic_cast<const SCEVNAryMulExpr *
>(lhsOperand);
1438 const auto rhsNAryMulExpr =
dynamic_cast<const SCEVNAryMulExpr *
>(rhsOperand);
1440 if (lhsNAryMulExpr && rhsNAryMulExpr)
1446 const auto lhsNAryAddExpr =
dynamic_cast<const SCEVNAryAddExpr *
>(lhsOperand);
1447 const auto rhsNAryAddExpr =
dynamic_cast<const SCEVNAryAddExpr *
>(rhsOperand);
1448 if ((lhsNAryMulExpr && rhsNAryAddExpr) || (rhsNAryMulExpr && lhsNAryAddExpr))
1452 const auto * mulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1453 auto * addExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1454 auto newAddExpr = SCEV::CloneAs<SCEVNAryExpr>(*addExpr);
1455 newAddExpr->AddOperand(mulExpr->Clone());
1456 return newAddExpr->Clone();
1459 const auto lhsInit =
dynamic_cast<const SCEVInit *
>(lhsOperand);
1460 const auto rhsInit =
dynamic_cast<const SCEVInit *
>(rhsOperand);
1461 if ((lhsNAryMulExpr && rhsInit) || (rhsNAryMulExpr && lhsInit))
1464 const auto * mulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1465 const auto * init = lhsInit ? lhsInit : rhsInit;
1469 const auto lhsConstant =
dynamic_cast<SCEVConstant *
>(lhsOperand);
1470 const auto rhsConstant =
dynamic_cast<SCEVConstant *
>(rhsOperand);
1475 const auto * mulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1476 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1480 if (lhsNAryMulExpr || rhsNAryMulExpr)
1483 const auto * mulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1484 return mulExpr->Clone();
1487 if (lhsInit && rhsInit)
1493 if ((lhsInit && rhsNAryAddExpr) || (rhsInit && lhsNAryAddExpr))
1496 const auto * init = lhsInit ? lhsInit : rhsInit;
1497 auto * nAryAddExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1498 auto newAddExpr = SCEV::CloneAs<SCEVNAryAddExpr>(*nAryAddExpr);
1499 newAddExpr->AddOperand(init->Clone());
1500 return newAddExpr->Clone();
1507 const auto * init = lhsInit ? lhsInit : rhsInit;
1508 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1512 if (lhsInit || rhsInit)
1515 const auto * init = lhsInit ? lhsInit : rhsInit;
1516 return init->Clone();
1519 if (lhsNAryAddExpr && rhsNAryAddExpr)
1522 auto lhsNewNAryAddExpr = SCEV::CloneAs<SCEVNAryAddExpr>(*lhsNAryAddExpr);
1523 for (
auto op : rhsNAryAddExpr->GetOperands())
1525 lhsNewNAryAddExpr->AddOperand(op->Clone());
1527 return lhsNewNAryAddExpr;
1534 auto * nAryAddExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1535 auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1536 auto newNAryAddExpr = SCEV::CloneAs<SCEVNAryAddExpr>(*nAryAddExpr);
1540 bool folded =
false;
1541 for (
size_t i = 0; i < newNAryAddExpr->NumOperands(); ++i)
1543 if (
auto existingConstant =
dynamic_cast<SCEVConstant *
>(newNAryAddExpr->GetOperands()[i]))
1546 auto foldedConstant =
ApplyAddFolding(existingConstant, constant, output);
1547 newNAryAddExpr->ReplaceOperand(i, foldedConstant);
1556 newNAryAddExpr->AddOperand(constant->Clone());
1559 return newNAryAddExpr;
1562 if (lhsNAryAddExpr || rhsNAryAddExpr)
1564 const auto * nAryAddExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1565 return nAryAddExpr->Clone();
1567 if (lhsConstant && rhsConstant)
1570 const auto lhsValue = lhsConstant->GetValue();
1571 const auto rhsValue = rhsConstant->GetValue();
1576 if (lhsConstant || rhsConstant)
1578 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1579 return constant->Clone();
1585std::unique_ptr<SCEVChainRecurrence>
1595 return SCEV::CloneAs<SCEVChainRecurrence>(*lhsChrec);
1597 return SCEV::CloneAs<SCEVChainRecurrence>(*rhsChrec);
1656 if (rhsSize > lhsSize)
1657 std::swap(lhsChrec, rhsChrec);
1659 std::unique_ptr<SCEVChainRecurrence> lhsStepRecurrence, rhsStepRecurrence;
1661 const auto lhsStep = *lhsChrec->
GetStep();
1665 throw std::logic_error(
"Could not get step for LHS in ComputeProductOfChrecs!");
1668 const auto rhsStep = *rhsChrec->
GetStep();
1671 throw std::logic_error(
"Could not get step for RHS in ComputeProductOfChrecs!");
1675 lhsStepRecurrence = SCEV::CloneAs<SCEVChainRecurrence>(*lhsStep);
1680 rhsStepRecurrence = SCEV::CloneAs<SCEVChainRecurrence>(*rhsStep);
1684 const auto rhsMarked = SCEV::CloneAs<SCEVChainRecurrence>(
1691 SCEV::CloneAs<SCEVChainRecurrence>(*
ApplyAddFolding(res1.get(), res2.get(), output));
1694 resFolded->AddOperandToFront(first);
1699std::unique_ptr<SCEV>
1709 if (
const auto *lhsUnknown =
dynamic_cast<const SCEVUnknown *
>(lhsOperand),
1710 *rhsUnknown =
dynamic_cast<const SCEVUnknown *
>(rhsOperand);
1711 lhsUnknown || rhsUnknown)
1718 if (lhsChrec && rhsChrec)
1720 if (&lhsChrec->GetLoop() != &rhsChrec->GetLoop())
1723 lhsChrec->GetLoop(),
1733 if (lhsChrec || rhsChrec)
1735 auto * chrec = lhsChrec ? lhsChrec : rhsChrec;
1736 auto * otherOperand = lhsChrec ? rhsOperand : lhsOperand;
1738 if (
auto constant =
dynamic_cast<const SCEVConstant *
>(otherOperand))
1740 if (constant->GetValue() == 1)
1743 return chrec->
Clone();
1746 if (constant->GetValue() == 0)
1753 const auto chrecOperands = chrec->GetOperands();
1755 for (
auto & operand : chrecOperands)
1763 const auto lhsNAryAddExpr =
dynamic_cast<const SCEVNAryAddExpr *
>(lhsOperand);
1764 const auto rhsNAryAddExpr =
dynamic_cast<const SCEVNAryAddExpr *
>(rhsOperand);
1765 if (lhsNAryAddExpr || rhsNAryAddExpr)
1769 const auto nAryAddExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1770 const auto other = lhsNAryAddExpr ? rhsOperand : lhsOperand;
1773 for (
auto operand : nAryAddExpr->GetOperands())
1776 resultAddExpr->AddOperand(std::move(product));
1778 return resultAddExpr;
1781 const auto lhsInit =
dynamic_cast<const SCEVInit *
>(lhsOperand);
1782 const auto rhsInit =
dynamic_cast<const SCEVInit *
>(rhsOperand);
1783 if (lhsInit && rhsInit)
1789 const auto lhsNAryMulExpr =
dynamic_cast<const SCEVNAryMulExpr *
>(lhsOperand);
1790 const auto rhsNAryMulExpr =
dynamic_cast<const SCEVNAryMulExpr *
>(rhsOperand);
1791 if ((lhsInit && rhsNAryMulExpr) || (rhsInit && lhsNAryMulExpr))
1794 const auto * init = lhsInit ? lhsInit : rhsInit;
1795 auto * nAryMulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1796 auto newNAryMulExpr = SCEV::CloneAs<SCEVNAryMulExpr>(*nAryMulExpr);
1797 newNAryMulExpr->AddOperand(init->Clone());
1798 return newNAryMulExpr->Clone();
1801 auto lhsConstant =
dynamic_cast<SCEVConstant *
>(lhsOperand);
1802 auto rhsConstant =
dynamic_cast<SCEVConstant *
>(rhsOperand);
1803 if ((lhsInit && rhsConstant && rhsConstant->GetValue() != 1)
1804 || (rhsInit && lhsConstant && lhsConstant->GetValue() != 1))
1807 const auto * init = lhsInit ? lhsInit : rhsInit;
1808 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1812 if (lhsInit || rhsInit)
1815 const auto * init = lhsInit ? lhsInit : rhsInit;
1816 return init->Clone();
1819 if (lhsNAryMulExpr && rhsNAryMulExpr)
1822 auto lhsNewNAryMulExpr = SCEV::CloneAs<SCEVNAryMulExpr>(*lhsNAryMulExpr);
1823 for (
auto op : rhsNAryMulExpr->GetOperands())
1825 lhsNewNAryMulExpr->AddOperand(op->Clone());
1827 return lhsNewNAryMulExpr;
1830 if ((lhsNAryMulExpr && rhsConstant && rhsConstant->GetValue() != 1)
1831 || (rhsNAryMulExpr && lhsConstant && lhsConstant->GetValue() != 1))
1834 auto * nAryMulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1835 auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1837 auto newNAryMulExpr = SCEV::CloneAs<SCEVNAryMulExpr>(*nAryMulExpr);
1839 bool folded =
false;
1840 for (
size_t i = 0; i < newNAryMulExpr->NumOperands(); ++i)
1842 if (
auto existingConstant =
dynamic_cast<SCEVConstant *
>(newNAryMulExpr->GetOperands()[i]))
1845 auto foldedConstant =
ApplyMulFolding(existingConstant, constant, output);
1846 newNAryMulExpr->ReplaceOperand(i, foldedConstant);
1855 newNAryMulExpr->AddOperand(constant->Clone());
1858 return newNAryMulExpr;
1861 if (lhsNAryMulExpr || rhsNAryMulExpr)
1863 const auto * nAryMulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1864 return nAryMulExpr->Clone();
1867 if (lhsConstant && rhsConstant)
1870 const auto lhsValue = lhsConstant->GetValue();
1871 const auto rhsValue = rhsConstant->GetValue();
1875 if (lhsConstant || rhsConstant)
1877 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1878 return constant->Clone();
1884std::unique_ptr<SCEV>
1888 if (
const auto c =
dynamic_cast<const SCEVConstant *
>(&scev))
1890 const auto value = c->GetValue();
1896 if (
const auto c =
dynamic_cast<const SCEVConstant *
>(
mul->GetLeftOperand());
1897 c && c->GetValue() == -1)
1899 return mul->GetRightOperand()->Clone();
1901 if (
const auto c =
dynamic_cast<const SCEVConstant *
>(
mul->GetRightOperand());
1904 return mul->GetLeftOperand()->Clone();
1920 auto deps = dependencyGraph[&output];
1921 if (deps.find(&output) != deps.end())
1923 if (deps[&output].count != 1)
1937 std::unordered_set<const rvsdg::Output *> visited{};
1938 std::unordered_set<const rvsdg::Output *> recursionStack{};
1947 std::unordered_set<const rvsdg::Output *> & visited,
1948 std::unordered_set<const rvsdg::Output *> & recursionStack)
1950 visited.insert(¤tOutput);
1951 recursionStack.insert(¤tOutput);
1953 for (
const auto & [depPtr, depCount] : dependencyGraph[¤tOutput])
1956 if (depPtr == ¤tOutput)
1961 if (depPtr == &originalOutput)
1965 if (visited.find(depPtr) != visited.end())
1973 recursionStack.erase(¤tOutput);
1989 if (
auto * binaryExpr =
dynamic_cast<const SCEVBinaryExpr *
>(&scev))
1991 return IsUnknown(*binaryExpr->GetLeftOperand()) ||
IsUnknown(*binaryExpr->GetLeftOperand());
1994 if (
auto * nAryExpr =
dynamic_cast<const SCEVNAryExpr *
>(&scev))
1996 for (
const auto operand : nAryExpr->GetOperands())
2004 throw std::logic_error(
"Invalid SCEV type in IsUnknown!\n");
2010 if (
typeid(a) !=
typeid(b))
2016 if (
auto * constantA =
dynamic_cast<const SCEVConstant *
>(&a))
2018 auto * constantB =
dynamic_cast<const SCEVConstant *
>(&b);
2019 return constantA->
GetValue() == constantB->GetValue();
2022 if (
auto * initA =
dynamic_cast<const SCEVInit *
>(&a))
2024 auto * initB =
dynamic_cast<const SCEVInit *
>(&b);
2028 if (
auto * binaryExprA =
dynamic_cast<const SCEVBinaryExpr *
>(&a))
2031 return StructurallyEqual(*binaryExprA->GetLeftOperand(), *binaryExprB->GetLeftOperand())
2032 &&
StructurallyEqual(*binaryExprA->GetRightOperand(), *binaryExprB->GetRightOperand());
2038 if (&chrecA->GetLoop() != &chrecB->GetLoop())
2040 if (&chrecA->GetOutput() != &chrecB->GetOutput())
2042 if (chrecA->NumOperands() != chrecB->NumOperands())
2044 for (
size_t i = 0; i < chrecA->NumOperands(); ++i)
2052 if (
auto * nAryExprA =
dynamic_cast<const SCEVNAryExpr *
>(&a))
2054 auto * nAryExprB =
dynamic_cast<const SCEVNAryExpr *
>(&b);
2055 if (nAryExprA->NumOperands() != nAryExprB->NumOperands())
2057 for (
size_t i = 0; i < nAryExprA->NumOperands(); ++i)
2059 if (!
StructurallyEqual(*nAryExprA->GetOperands()[i], *nAryExprB->GetOperands()[i]))
static rvsdg::Input & getAddressInput(const rvsdg::Node &node) noexcept
static std::unique_ptr< SCEVAddExpr > Create(std::unique_ptr< SCEV > left, std::unique_ptr< SCEV > right)
rvsdg::ThetaNode & GetLoop() const noexcept
static std::unique_ptr< SCEVChainRecurrence > Create(rvsdg::ThetaNode &loop, rvsdg::Output &output)
static bool IsQuadratic(const SCEVChainRecurrence &chrec)
static bool IsConstant(const SCEVChainRecurrence &chrec)
std::optional< std::unique_ptr< SCEV > > GetStep() const
SCEV * GetStartValue() const
static bool IsAffine(const SCEVChainRecurrence &chrec)
static std::unique_ptr< SCEVConstant > Create(const int64_t value)
static bool IsNonZero(const SCEVConstant *c)
static std::unique_ptr< SCEVInit > Create(rvsdg::Output &prePointer)
rvsdg::Output & GetPrePointer() const noexcept
static std::unique_ptr< SCEVMulExpr > Create(std::unique_ptr< SCEV > left, std::unique_ptr< SCEV > right)
static std::unique_ptr< SCEVNAryAddExpr > Create(Args &&... operands)
size_t NumOperands() const
void RemoveOperand(const size_t index)
SCEV * GetOperand(const size_t index) const
void ReplaceOperand(const size_t index, const std::unique_ptr< SCEV > &operand)
std::vector< SCEV * > GetOperands() const
static std::unique_ptr< SCEVNAryMulExpr > Create(Args &&... operands)
static std::unique_ptr< SCEVPlaceholder > Create(rvsdg::Output &PrePointer_)
static std::unique_ptr< SCEVUnknown > Create()
virtual std::unique_ptr< SCEV > Clone() const =0
size_t GetTripCount(const rvsdg::ThetaNode &thetaNode) const
std::unordered_map< const rvsdg::ThetaNode *, size_t > TripCountMap_
const std::unordered_map< rvsdg::Output *, std::unique_ptr< SCEV > > & GetSCEVMap() const noexcept
size_t GetNumLoops() const
size_t GetNumTotalLoopVars() const
std::unordered_set< const rvsdg::Output * > LoopVars_
void SetTripCount(const rvsdg::ThetaNode &thetaNode, const size_t tripCount)
const std::unordered_map< const rvsdg::ThetaNode *, size_t > & GetTripCountMap() const noexcept
void InsertSCEV(rvsdg::Output &output, const std::unique_ptr< SCEV > &scev)
std::unordered_map< rvsdg::Output *, std::unique_ptr< SCEVChainRecurrence > > ChrecMap_
std::unique_ptr< SCEVChainRecurrence > TryGetChrecForOutput(rvsdg::Output &output) const
const std::unordered_map< rvsdg::Output *, std::unique_ptr< SCEVChainRecurrence > > & GetChrecMap() const noexcept
Context(const Context &)=delete
void AddLoopVar(const rvsdg::Output &var)
std::unordered_map< rvsdg::Output *, std::unique_ptr< SCEV > > SCEVMap_
size_t GetNumTotalInductionVariables() const
Context(Context &&)=delete
Context & operator=(Context &&)=delete
Context & operator=(const Context &)=delete
std::unique_ptr< SCEV > TryGetSCEVForOutput(rvsdg::Output &output) const
void InsertChrec(rvsdg::Output &output, const std::unique_ptr< SCEVChainRecurrence > &chrec)
static std::unique_ptr< Context > Create()
int GetNumInductionVariablesWithOrder(const size_t n) const
~Statistics() noexcept override=default
static std::string GetTripCountString(const std::unordered_map< const rvsdg::ThetaNode *, size_t > &tripCountMap)
static std::unique_ptr< Statistics > Create(const util::FilePath &sourceFile)
void Stop(const Context &context) noexcept
std::unordered_map< rvsdg::Output *, DependencyInfo > DependencyMap
std::unordered_map< rvsdg::Output *, DependencyMap > DependencyGraph
static std::unique_ptr< SCEVChainRecurrence > ComputeProductOfChrecs(SCEVChainRecurrence *lhsChrec, SCEVChainRecurrence *rhsChrec, rvsdg::Output &output)
std::optional< std::unique_ptr< SCEV > > TryReplaceInitForSCEV(const SCEV &scev, rvsdg::Output &output)
std::unique_ptr< SCEVChainRecurrence > GetOrCreateStepForSCEV(rvsdg::Output &output, const SCEV &scevTree, rvsdg::ThetaNode &thetaNode)
void PerformSCEVAnalysis(rvsdg::ThetaNode &thetaNode)
void Run(rvsdg::RvsdgModule &rvsdgModule, util::StatisticsCollector &statisticsCollector) override
Perform RVSDG transformation.
static void FindDependenciesForSCEV(const SCEV &scev, DependencyMap &dependencies, DependencyOp op)
static bool IsStepZero(const SCEV &stepSCEV)
static std::unique_ptr< SCEV > ApplyMulFolding(SCEV *lhsOperand, SCEV *rhsOperand, rvsdg::Output &output)
Apply folding rules for multiplication to combine two SCEV operands into one.
static std::unique_ptr< SCEV > FoldNAryExpression(SCEVNAryExpr &expression, rvsdg::Output &output)
Try to combine the constants in an n-ary expression (Add or Mul) into themselves.
static std::optional< size_t > SolveQuadraticEquation(int64_t a, int64_t b, int64_t c)
Tries to find a solution to the quadratic equation a^2 x + b x + c = 0 using integer arithmetic.
DependencyGraph CreateDependencyGraph(const rvsdg::ThetaNode &thetaNode) const
static std::unique_ptr< SCEV > ApplyAddFolding(SCEV *lhsOperand, SCEV *rhsOperand, rvsdg::Output &output)
Apply folding rules for addition to combine two SCEV operands into one.
static bool HasCycleThroughOthers(rvsdg::Output ¤tOutput, const rvsdg::Output &originalOutput, DependencyGraph &dependencyGraph, std::unordered_set< const rvsdg::Output * > &visited, std::unordered_set< const rvsdg::Output * > &recursionStack)
~ScalarEvolution() noexcept override
std::unique_ptr< Context > Context_
static bool IsStepPositive(const SCEV &stepSCEV)
void CombineChrecsAcrossLoops()
static bool CanCreateChainRecurrence(rvsdg::Output &output, DependencyGraph &dependencyGraph)
std::unique_ptr< SCEV > ComputeSCEVForGepInnerOffset(const rvsdg::SimpleNode &gepNode, size_t inputIndex, const rvsdg::Type &type)
std::optional< size_t > GetPredictedTripCount(rvsdg::ThetaNode &thetaNode)
static bool StructurallyEqual(const SCEV &a, const SCEV &b)
static bool IsUnknown(const SCEV &scev)
std::unordered_map< const rvsdg::Output *, std::unique_ptr< SCEV > > GetSCEVMap() const
static bool IsStepNegative(const SCEV &stepSCEV)
static std::vector< rvsdg::Output * > TopologicalSort(DependencyGraph &dependencyGraph)
std::unique_ptr< SCEV > GetOrCreateSCEVForOutput(rvsdg::Output &output)
std::unique_ptr< SCEVChainRecurrence > GetOrCreateChainRecurrence(rvsdg::Output &output, const SCEV &scev, rvsdg::ThetaNode &thetaNode)
std::unordered_map< const rvsdg::ThetaNode *, size_t > GetTripCountMap() const noexcept
static std::unique_ptr< SCEV > GetNegativeSCEV(const SCEV &scev)
void AnalyzeRegion(rvsdg::Region ®ion)
static std::optional< size_t > ComputeBackedgeTakenCountForChrec(const SCEVChainRecurrence &chrec, int64_t bound, const rvsdg::SimpleOperation *comparisonOperation)
Region & GetRootRegion() const noexcept
size_t ninputs() const noexcept
Represent acyclic RVSDG subgraphs.
NodeRange Nodes() noexcept
const std::optional< util::FilePath > & SourceFilePath() const noexcept
NodeInput * input(size_t index) const noexcept
LoopVar MapPreLoopVar(const rvsdg::Output &argument) const
Maps variable at start of loop iteration to full varibale description.
std::vector< LoopVar > GetLoopVars() const
Returns all loop variables.
RegionResult * predicate() const noexcept
bool insert(ItemType item)
void CollectDemandedStatistics(std::unique_ptr< Statistics > statistics)
util::Timer & GetTimer(const std::string &name)
util::Timer & AddTimer(std::string name)
void AddMeasurement(std::string name, T value)
Global memory state passed between functions.
size_t GetTypeAllocSize(const rvsdg::Type &type)
static util::StatisticsCollector statisticsCollector
std::optional< int64_t > tryGetConstantSignedInteger(const rvsdg::Output &output)
Output & traceOutput(Output &output, bool mayEnterSubregions, const Region *withinRegion)
static bool ThetaLoopVarIsInvariant(const ThetaNode::LoopVar &loopVar) noexcept
@ State
Designate a state type.
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.
rvsdg::Input * input
Variable at loop entry (input to theta).
rvsdg::Input * post
Variable after iteration (output result from subregion).