22#include <unordered_map>
78 static std::unique_ptr<Statistics>
81 return std::make_unique<Statistics>(sourceFile);
167 throw std::logic_error(
"Output does not belong to a congruence set");
193 auto nextSet =
sets_.size();
199 if (
sets_[it->second].leader == &leader)
205 sets_[it->second].followers.Remove(&leader);
206 it->second = nextSet;
210 sets_.emplace_back(leader);
223 return *
sets_[index].leader;
239 const bool newFollower =
sets_[index].followers.insert(&follower);
252 if (
sets_[it->second].leader == &follower)
253 throw std::logic_error(
"Cannot turn a leader into a follower");
255 const bool removed =
sets_[it->second].followers.Remove(&follower);
271 return sets_[index].followers;
299 const auto o1Set = context.
getSetFor(o1);
300 const auto o2Set = context.
getSetFor(o2);
302 return o1Set == o2Set;
325 if (!simpleNode1 || !simpleNode2)
328 if (simpleNode1->ninputs() != simpleNode2->ninputs())
331 if (simpleNode1->GetOperation() != simpleNode2->GetOperation())
336 for (
auto & input : simpleNode1->Inputs())
338 const auto origin1 = input.origin();
339 const auto origin2 = simpleNode2->
input(input.index())->origin();
355 for (
auto & output : leader.
Outputs())
380 for (
size_t i = 0; i < leader.
noutputs(); i++)
382 const auto & leaderOutput = *leader.
output(i);
383 const auto & followerOutput = *follower.
output(i);
384 const auto leaderSet = context.
getSetFor(leaderOutput);
411 const auto & output0Leader = context.
getLeader(output0Set);
437 if (existingLeaderNode == &node)
439 leaders.push_back(&node);
456 for (
auto leader : leaders)
467 leaders.push_back(&node);
488 if (leaderNode == &node)
501 const auto tryFindCongruentUserOf = [&](
const rvsdg::Output & output) ->
bool
505 for (
auto & user : output.Users())
507 if (user.index() != 0)
515 if (otherNode == &node)
520 if (otherNode != otherNodeLeader)
537 const auto origin0Set = context.
getSetFor(*origin);
538 const auto & origin0Leader = context.
getLeader(origin0Set);
539 const auto & origin0Followers = context.
getFollowers(origin0Set);
540 if (tryFindCongruentUserOf(origin0Leader))
542 for (
auto follower : origin0Followers.Items())
544 if (tryFindCongruentUserOf(*follower))
587 const std::vector<CommonNodeElimination::Context::CongruenceSetIndex> & partitions,
604 const auto currentPartition = context.
tryGetSetFor(*argument);
605 const auto key = std::make_pair(currentPartition, partitions[argument->index()]);
610 if (
const auto it = newSets.find(key); it != newSets.end())
613 const auto toFollow = it->second;
616 if (currentPartition == toFollow)
655 bool anyChanges =
false;
659 if (subregion.narguments() == 0)
666 std::vector<size_t> partitions(subregion.narguments());
671 for (
const auto argument : subregion.Arguments())
673 if (
const auto input = argument->input())
676 partitions[argument->index()] = context.
getSetFor(*input->origin());
681 partitions[argument->index()] = nextUniquePartitionKey++;
706static std::optional<CommonNodeElimination::Context::CongruenceSetIndex>
711 std::optional<CommonNodeElimination::Context::CongruenceSetIndex> sharedCongruenceSet;
717 const auto inputCongruenceSet = context.
getSetFor(*argument->input()->origin());
718 if (!sharedCongruenceSet.has_value())
720 sharedCongruenceSet = inputCongruenceSet;
722 else if (*sharedCongruenceSet != inputCongruenceSet)
735 return sharedCongruenceSet;
747[[nodiscard]]
static size_t
755 const auto set = context.
getSetFor(*branchResult->origin());
773[[nodiscard]]
static bool
783 for (
size_t i = 0; i < numResults; i++)
787 if (firstSet != secondSet)
806 std::unordered_map<size_t, CommonNodeElimination::Context::CongruenceSetIndex> & leaderHashes,
814 const auto [_, inserted] = leaderHashes.emplace(hash, congruenceSet);
840 std::unordered_map<size_t, CommonNodeElimination::Context::CongruenceSetIndex> & leaderHashes,
847 const auto [it, inserted] = leaderHashes.emplace(hash, 0);
855 auto & otherLeader = context.
getLeader(it->second);
883 std::unordered_map<size_t, CommonNodeElimination::Context::CongruenceSetIndex> leaderHashes;
889 bool skipInvarianceCheck =
false;
892 const auto existingSet = context.
tryGetSetFor(*exitVar.output);
898 const auto & exisitingLeader = context.
getLeader(existingSet);
899 if (&exisitingLeader == exitVar.output)
920 skipInvarianceCheck =
true;
924 if (!skipInvarianceCheck)
928 entryVarCongruenceSet.has_value())
930 context.
addFollower(*entryVarCongruenceSet, *exitVar.output);
961 bool onlyUpdateLoopInvariants = !anyChanges;
967 std::vector<CommonNodeElimination::Context::CongruenceSetIndex> partitions;
968 for (
const auto & loopVar : loopVars)
970 partitions.push_back(context.
getSetFor(*loopVar.post->origin()));
984 const auto isOutputInvariant = [&](
const auto & loopVar) ->
rvsdg::Output *
989 return loopVar.input->origin();
992 const auto resultSet = context.
getSetFor(*loopVar.post->origin());
993 const auto & resultSetLeader = context.
getLeader(resultSet);
995 if (&resultSetLeader == loopVar.pre)
997 return loopVar.input->origin();
1004 const auto otherLoopVar = theta.
MapPreLoopVar(resultSetLeader);
1012 const auto otherLoopVarPostSet = context.
getSetFor(*otherLoopVar.post->origin());
1013 const auto & otherLoopVarPostLeader = context.
getLeader(otherLoopVarPostSet);
1014 if (&otherLoopVarPostLeader == otherLoopVar.pre)
1016 return otherLoopVar.input->origin();
1027 resultToOutputSetMapping;
1028 for (
auto & loopVar : loopVars)
1032 if (
const auto origin = isOutputInvariant(loopVar))
1034 auto inputCongruenceSet = context.
getSetFor(*origin);
1035 context.
addFollower(inputCongruenceSet, *loopVar.output);
1039 if (onlyUpdateLoopInvariants)
1043 const auto resultSet = context.
getSetFor(*loopVar.post->origin());
1044 const auto it = resultToOutputSetMapping.find(resultSet);
1045 if (it != resultToOutputSetMapping.end())
1052 resultToOutputSetMapping.emplace(resultSet, outputSet);
1116 if (node->ninputs() == 0)
1118 markSimpleTopNode(simple, leaders, context);
1122 markSimpleNode(simple, context);
1137 const auto outputSet = context.
getSetFor(output);
1139 auto & leader = context.
getLeader(outputSet);
1140 if (&leader == &output)
1152 bool divertInSubregions =
false;
1157 divertInSubregions =
true;
1161 divertInSubregions =
true;
1165 divertInSubregions =
true;
1169 divertInSubregions =
true;
1177 if (divertInSubregions)
1190 for (
auto argument : region.
Arguments())
1207 for (
auto & output : node->Outputs())
1213 region.
prune(
false);
1220 rvsdg::RvsdgModule & module,
1223 auto & rvsdg =
module.Rvsdg();
1224 auto & rootRegion = rvsdg.GetRootRegion();
1227 auto statistics = Statistics::Create(module.SourceFilePath().value());
1229 statistics->startMarkStatistics(rvsdg);
1232 statistics->endMarkStatistics();
1234 statistics->startDivertStatistics();
1236 statistics->endDivertStatistics(rvsdg);
1246 statistics->startPruneStatistics();
1248 statistics->stopPruneStatistics();
static constexpr auto NoCongruenceSetIndex
CongruenceSetIndex getOrCreateSetForLeader(const rvsdg::Output &leader)
const util::HashSet< const rvsdg::Output * > & getFollowers(CongruenceSetIndex index) const
CongruenceSetIndex numCongruenceSets() const
void addFollower(CongruenceSetIndex index, const rvsdg::Output &follower)
size_t CongruenceSetIndex
bool hasSet(const rvsdg::Output &output) const
CongruenceSetIndex getSetFor(const rvsdg::Output &output) const
CongruenceSetIndex tryGetSetFor(const rvsdg::Output &output) const
const rvsdg::Output & getLeader(CongruenceSetIndex index) const
std::vector< CongruenceSet > sets_
std::unordered_map< const rvsdg::Output *, CongruenceSetIndex > congruenceSetMapping_
void startMarkStatistics(const rvsdg::Graph &graph) noexcept
void startDivertStatistics() noexcept
void stopPruneStatistics() noexcept
void endMarkStatistics() noexcept
void endDivertStatistics(const rvsdg::Graph &graph) noexcept
const char * PruneTimerLabel_
void startPruneStatistics() noexcept
~Statistics() override=default
const char * DivertTimerLabel_
Statistics(const util::FilePath &sourceFile)
static std::unique_ptr< Statistics > Create(const util::FilePath &sourceFile)
const char * MarkTimerLabel_
Common Node Elimination Discovers simple nodes, region arguments and structural node outputs that are...
~CommonNodeElimination() noexcept override
Conditional operator / pattern matching.
ExitVar MapOutputExitVar(const rvsdg::Output &output) const
Maps gamma output to exit variable description.
std::vector< ExitVar > GetExitVars() const
Gets all exit variables for this gamma.
NodeOutput * output(size_t index) const noexcept
OutputIteratorRange Outputs() noexcept
rvsdg::Region * region() const noexcept
size_t ninputs() const noexcept
size_t noutputs() const noexcept
void divert_users(jlm::rvsdg::Output *new_origin)
const std::shared_ptr< const rvsdg::Type > & Type() const noexcept
A phi node represents the fixpoint of mutually recursive definitions.
Represents the argument of a region.
Represent acyclic RVSDG subgraphs.
RegionArgumentRange Arguments() noexcept
void prune(bool recursive)
size_t narguments() const noexcept
NodeInput * input(size_t index) const noexcept
NodeOutput * output(size_t index) const noexcept
SubregionIteratorRange Subregions()
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.
rvsdg::Region * subregion() const noexcept
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.
static void lookupOrInsertGammaExitVarInHashmap(const rvsdg::GammaNode::ExitVar &exitVar, const rvsdg::GammaNode &gamma, std::unordered_map< size_t, CommonNodeElimination::Context::CongruenceSetIndex > &leaderHashes, CommonNodeElimination::Context &context)
static bool partitionArguments(const rvsdg::Region ®ion, const std::vector< CommonNodeElimination::Context::CongruenceSetIndex > &partitions, CommonNodeElimination::Context &context)
static util::StatisticsCollector statisticsCollector
static void markSimpleTopNode(const rvsdg::SimpleNode &node, TopNodeLeaderList &leaders, CommonNodeElimination::Context &context)
static bool markSubregionsFromInputs(const rvsdg::StructuralNode &node, CommonNodeElimination::Context &context)
static void markGamma(const rvsdg::GammaNode &gamma, CommonNodeElimination::Context &context)
static void markGraphImports(const rvsdg::Region ®ion, CommonNodeElimination::Context &context)
static bool checkNodesCongruent(const rvsdg::Node &node1, const rvsdg::Node &node2, CommonNodeElimination::Context &context)
static void divertOutput(rvsdg::Output &output, CommonNodeElimination::Context &context)
static void divertInRegion(rvsdg::Region &, CommonNodeElimination::Context &)
void markNodeAsLeader(const rvsdg::Node &leader, CommonNodeElimination::Context &context)
static void insertGammaExitVarInHashmap(rvsdg::GammaNode::ExitVar &exitVar, CommonNodeElimination::Context::CongruenceSetIndex congruenceSet, std::unordered_map< size_t, CommonNodeElimination::Context::CongruenceSetIndex > &leaderHashes, CommonNodeElimination::Context &context)
static size_t getGammaExitVariableHash(const rvsdg::GammaNode::ExitVar &exitVar, CommonNodeElimination::Context &context)
std::vector< const rvsdg::Node * > TopNodeLeaderList
static void markRegion(const rvsdg::Region &, CommonNodeElimination::Context &context)
static std::optional< CommonNodeElimination::Context::CongruenceSetIndex > tryGetGammaExitVarCongruenceSet(rvsdg::GammaNode::ExitVar &exitVar, CommonNodeElimination::Context &context)
static bool areOutputsCongruent(const rvsdg::Output &o1, const rvsdg::Output &o2, CommonNodeElimination::Context &context)
void markNodesAsCongruent(const rvsdg::Node &leader, const rvsdg::Node &follower, CommonNodeElimination::Context &context)
static void divertInStructuralNode(rvsdg::StructuralNode &node, CommonNodeElimination::Context &context)
static void markTheta(const rvsdg::ThetaNode &theta, CommonNodeElimination::Context &context)
static bool areGammaExitVariablesCongruent(const rvsdg::GammaNode::ExitVar &first, const rvsdg::GammaNode::ExitVar &second, CommonNodeElimination::Context &context)
static void markStructuralNode(const rvsdg::StructuralNode &node, CommonNodeElimination::Context &context)
static void markSimpleNode(const rvsdg::SimpleNode &node, CommonNodeElimination::Context &context)
const rvsdg::SimpleNode * tryGetLeaderNode(const rvsdg::SimpleNode &node, CommonNodeElimination::Context &context)
void MatchTypeOrFail(T &obj, const Fns &... fns)
Pattern match over subclass type of given object.
static bool ThetaLoopVarIsInvariant(const ThetaNode::LoopVar &loopVar) noexcept
void MatchType(T &obj, const Fns &... fns)
Pattern match over subclass type of given object.
detail::TopDownTraverserGeneric< true > TopDownConstTraverser
Traverser for visiting every node in a const region in a top down order.
size_t ninputs(const rvsdg::Region *region) noexcept
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.
detail::TopDownTraverserGeneric< false > TopDownTraverser
Traverser for visiting every node in a region in a top down order.
std::size_t CombineHashes(std::size_t hash, Args... args)
CongruenceSet(const rvsdg::Output &leader)
util::HashSet< const rvsdg::Output * > followers
const rvsdg::Output * leader
A variable routed out of all gamma regions as result.
rvsdg::Output * output
Output of gamma.
std::vector< rvsdg::Input * > branchResult
Variable exit points (results per subregion).
rvsdg::Input * input
Variable at loop entry (input to theta).