Jlm
Loading...
Searching...
No Matches
NodeHoisting.cpp
Go to the documentation of this file.
1/*
2 * Copyright 2017 Nico Reißmann <nico.reissmann@gmail.com>
3 * See COPYING for terms of redistribution.
4 */
5
15#include <jlm/rvsdg/control.hpp>
16#include <jlm/rvsdg/gamma.hpp>
18#include <jlm/rvsdg/Phi.hpp>
19#include <jlm/rvsdg/theta.hpp>
22#include <jlm/util/time.hpp>
23
24#include <algorithm>
25#include <deque>
26
27namespace jlm::llvm
28{
29
31{
32public:
33 ~Statistics() override = default;
34
35 explicit Statistics(const util::FilePath & sourceFile)
36 : util::Statistics(Statistics::Id::PushNodes, sourceFile)
37 {}
38
39 void
40 start(const rvsdg::Graph & graph) noexcept
41 {
42 AddMeasurement(Label::NumRvsdgInputsBefore, jlm::rvsdg::ninputs(&graph.GetRootRegion()));
43 AddTimer(Label::Timer).start();
44 }
45
46 void
47 end(const rvsdg::Graph & graph) noexcept
48 {
49 AddMeasurement(Label::NumRvsdgInputsAfter, jlm::rvsdg::ninputs(&graph.GetRootRegion()));
50 GetTimer(Label::Timer).stop();
51 }
52
53 static std::unique_ptr<Statistics>
54 Create(const util::FilePath & sourceFile)
55 {
56 return std::make_unique<Statistics>(sourceFile);
57 }
58};
59
61{
62public:
63 explicit Context(rvsdg::LambdaNode & lambdaNode)
64 : LambdaSubregion_(lambdaNode.subregion())
65 {}
66
68 getLambdaSubregion() const noexcept
69 {
70 return *LambdaSubregion_;
71 }
72
73 void
74 addTargetRegion(const rvsdg::Node & node, rvsdg::Region & region) noexcept
75 {
76 JLM_ASSERT(TargetRegion_.find(&node) == TargetRegion_.end());
77 TargetRegion_[&node] = &region;
78 }
79
81 getTargetRegion(const rvsdg::Node & node) const noexcept
82 {
83 return *TargetRegion_.at(&node);
84 }
85
86 static std::unique_ptr<Context>
88 {
89 return std::make_unique<Context>(lambdaNode);
90 }
91
92private:
94 std::unordered_map<const rvsdg::Node *, rvsdg::Region *> TargetRegion_{};
95};
96
97NodeHoisting::~NodeHoisting() noexcept = default;
98
100 : Transformation("NodeHoisting")
101{}
102
103bool
105{
106 if (!is<MemoryStateType>(loopVar.output->Type()))
107 return false;
108
109 if (loopVar.pre->nusers() != 1)
110 return false;
111
112 // FIXME: This check fails if we have a simple node followed for example by a gamma node.
113 // The consequence is that nodes are not pushed out as much as they could.
114 const auto userNode = rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(*loopVar.pre->Users().begin());
115 const auto originNode = rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(*loopVar.post->origin());
116
117 if (userNode == nullptr || originNode == nullptr || userNode != originNode)
118 return false;
119
120 return true;
121}
122
125{
126 // Handle lambda region arguments
128 {
129 return *output.region();
130 }
131
132 // Handle gamma region arguments
133 if (const auto gammaNode = rvsdg::TryGetRegionParentNode<rvsdg::GammaNode>(output))
134 {
135 if (is<IOStateType>(output.Type()))
136 {
137 // Do not hoist nodes with IO state edges out of gamma nodes.
138 return *output.region();
139 }
140
141 const auto roleVar = gammaNode->MapBranchArgument(output);
142 if (const auto entryVar = std::get_if<rvsdg::GammaNode::EntryVar>(&roleVar))
143 {
144 return computeTargetRegion(*entryVar->input->origin());
145 }
146
147 return *output.region();
148 }
149
150 // Handle theta region arguments
151 if (const auto thetaNode = rvsdg::TryGetRegionParentNode<rvsdg::ThetaNode>(output))
152 {
153 const auto loopVar = thetaNode->MapPreLoopVar(output);
155 {
156 return computeTargetRegion(*loopVar.input->origin());
157 }
158
160 {
161 return computeTargetRegion(*loopVar.input->origin());
162 }
163
164 return *output.region();
165 }
166
167 // Handle gamma outputs
168 if (const auto gammaNode = rvsdg::TryGetOwnerNode<rvsdg::GammaNode>(output))
169 {
170 return context_->getTargetRegion(*gammaNode);
171 }
172
173 // Handle theta outputs
174 if (const auto thetaNode = rvsdg::TryGetOwnerNode<rvsdg::ThetaNode>(output))
175 {
176 return context_->getTargetRegion(*thetaNode);
177 }
178
179 // Handle simple node outputs
180 if (const auto node = rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(output))
181 {
182 return context_->getTargetRegion(*node);
183 }
184
185 throw std::logic_error("Unhandled output type!");
186}
187
188static bool
190{
191 for (auto & input : node.Inputs())
192 {
193 if (input.Type()->Kind() != rvsdg::TypeKind::Value)
194 return false;
195 }
196
197 return true;
198}
199
200static rvsdg::Region &
201limitTargetRegion(const rvsdg::Node & node, rvsdg::Region & targetRegion)
202{
203 JLM_ASSERT(node.region() != &targetRegion);
204
205 if (hasOnlyValueInputs(node))
206 {
207 // Pure nodes can be hoisted out of gamma and theta nodes.
208 return targetRegion;
209 }
210
211 if (is<LoadNonVolatileOperation>(node.GetOperation()))
212 {
213 // LoadNonVolatileOperation nodes can also be hoisted out of gamma and theta nodes.
214 return targetRegion;
215 }
216
217 // For all other nodes, we want to limit the target region to the lowest gamma node.
218 auto currentRegion = node.region();
219 do
220 {
221 if (dynamic_cast<rvsdg::GammaNode *>(currentRegion->node()))
222 {
223 break;
224 }
225
226 currentRegion = currentRegion->node()->region();
227 } while (currentRegion != &targetRegion);
228
229 return *currentRegion;
230}
231
234{
235 if (node.ninputs() == 0)
236 {
237 // All nullary operations to date have exactly one output
238 JLM_ASSERT(node.noutputs() == 1);
239
240 // Control constants are used to instruct the control flow graph creation,
241 // and will be removed in the back-end, so there is no need to hoist them.
242 const auto outputType = node.output(0)->Type();
243 if (is<rvsdg::ControlType>(outputType))
244 return *node.region();
245
246 // Other constants should be moved to the top-level of the function
247 return context_->getLambdaSubregion();
248 }
249
250 // Compute target regions for all the inputs of the node
251 rvsdg::Region * greatestCommonTargetRegion = nullptr;
252
253 for (auto & input : node.Inputs())
254 {
255 auto & targetRegion = computeTargetRegion(*input.origin());
256 if (&targetRegion == node.region())
257 {
258 // One of the node's predecessors cannot be hoisted, which means we can also not hoist this
259 // node
260 return *node.region();
261 }
262
263 // If we already have a common target region that is lower, keep it
264 if (greatestCommonTargetRegion
265 && greatestCommonTargetRegion->getDepth() >= targetRegion.getDepth())
266 continue;
267 greatestCommonTargetRegion = &targetRegion;
268 }
269
270 greatestCommonTargetRegion = &limitTargetRegion(node, *greatestCommonTargetRegion);
271
272 // Return the lowest-most common target region in the region tree among all inputs
273 JLM_ASSERT(greatestCommonTargetRegion);
274 return *greatestCommonTargetRegion;
275}
276
277void
279{
280 for (const auto node : rvsdg::TopDownConstTraverser(&region))
281 {
283 *node,
284 [&](const rvsdg::StructuralNode & structuralNode)
285 {
286 // FIXME: We currently do not allow structural nodes (gamma and theta nodes) to be hoisted
287 context_->addTargetRegion(structuralNode, *structuralNode.region());
288
289 // Handle innermost regions
290 for (auto & subregion : structuralNode.Subregions())
291 {
292 markNodes(subregion);
293 }
294 },
295 [&](const rvsdg::SimpleNode & simpleNode)
296 {
297 rvsdg::Region & targetRegion = computeTargetRegion(simpleNode);
298 context_->addTargetRegion(*node, targetRegion);
299 },
300 []()
301 {
302 throw std::logic_error("Unhandled node type!");
303 });
304 }
305}
306
309{
310 if (input.region() == &targetRegion)
311 return input;
312
313 const auto & operand = *input.origin();
314
315 // Handle gamma subregion arguments
316 if (const auto gammaNode = rvsdg::TryGetRegionParentNode<rvsdg::GammaNode>(operand))
317 {
318 const auto roleVar = gammaNode->MapBranchArgument(operand);
319 if (const auto entryVar = std::get_if<rvsdg::GammaNode::EntryVar>(&roleVar))
320 {
321 return getUserFromTargetRegion(*entryVar->input, targetRegion);
322 }
323 }
324
325 // Handle theta subregion arguments
326 if (const auto thetaNode = rvsdg::TryGetRegionParentNode<rvsdg::ThetaNode>(operand))
327 {
328 const auto loopVar = thetaNode->MapPreLoopVar(operand);
330 return getUserFromTargetRegion(*loopVar.input, targetRegion);
331 }
332
333 if (const auto simpleNode = rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(operand))
334 {
335 if (is<LoadNonVolatileOperation>(simpleNode->GetOperation()))
336 {
337 auto & memStateInput = LoadNonVolatileOperation::MapMemoryStateOutputToInput(operand);
338 return getUserFromTargetRegion(memStateInput, targetRegion);
339 }
340 }
341
342 throw std::logic_error("Unhandled output type!");
343}
344
345std::vector<rvsdg::Input *>
347{
348 std::vector<rvsdg::Input *> users;
349 for (auto & input : node.Inputs())
350 {
351 auto & user = getUserFromTargetRegion(input, targetRegion);
352 users.push_back(&user);
353 }
354
355 return users;
356}
357
358static rvsdg::Input *
360{
361 JLM_ASSERT(output.Type()->Kind() == rvsdg::TypeKind::State);
362
363 const auto simpleNode = rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(output);
364 JLM_ASSERT(simpleNode);
365
367 simpleNode->GetOperation(),
368 [&output](const LoadNonVolatileOperation &)
369 {
370 return &LoadOperation::MapMemoryStateOutputToInput(output);
371 },
372 [&output](const StoreNonVolatileOperation &)
373 {
374 return &StoreOperation::MapMemoryStateOutputToInput(output);
375 },
376 [&output](const CallOperation &)
377 {
378 if (is<IOStateType>(output.Type()))
379 {
380 return &CallOperation::mapIOStateOutputToInput(output);
381 }
382
383 if (is<MemoryStateType>(output.Type()))
384 {
385 return &CallOperation::mapMemoryStateOutputToInput(output);
386 }
387
388 throw std::logic_error(
389 util::strfmt("Unhandled output type: ", output.Type()->debug_string()));
390 },
391 [](const VariadicArgumentListOperation &) -> rvsdg::Input *
392 {
393 return nullptr;
394 },
395 [](const CallEntryMemoryStateMergeOperation &) -> rvsdg::Input *
396 {
397 return nullptr;
398 },
399 [](const AllocaOperation &) -> rvsdg::Input *
400 {
401 return nullptr;
402 },
403 [&simpleNode]() -> rvsdg::Input *
404 {
405 throw std::logic_error(
406 util::strfmt("Unhandled operation type: ", simpleNode->DebugString()));
407 });
408}
409
410void
411NodeHoisting::copyNodeToTargetRegion(rvsdg::Node & node) const
412{
413 auto & targetRegion = context_->getTargetRegion(node);
414 JLM_ASSERT(&targetRegion != node.region());
415
416 const auto users = getUsersFromTargetRegion(node, targetRegion);
417
418 std::vector<rvsdg::Output *> operands;
419 operands.reserve(users.size());
420 std::transform(
421 users.begin(),
422 users.end(),
423 std::back_inserter(operands),
424 [](const rvsdg::Input * input) -> rvsdg::Output *
425 {
426 return input->origin();
427 });
428
429 const auto copiedNode = node.copy(&targetRegion, operands);
430
431 // FIXME: I really would like to have a zip function here, but C++ does not really seem to have
432 // anything better to offer
433 auto itOrg = std::begin(node.Outputs());
434 const auto endOrg = std::end(node.Outputs());
435 auto itCpy = std::begin(copiedNode->Outputs());
436 const auto endCpy = std::end(copiedNode->Outputs());
437 JLM_ASSERT(std::distance(itOrg, endOrg) == std::distance(itCpy, endCpy));
438
439 for (; itOrg != endOrg; ++itOrg, ++itCpy)
440 {
441 auto & outputOrg = *itOrg;
442 auto & outputCpy = *itCpy;
443
444 if (outputOrg.Type()->Kind() == rvsdg::TypeKind::State)
445 {
446 if (auto inputOrg = mapStateOutputToInput(outputOrg))
447 {
448 outputOrg.divert_users(inputOrg->origin());
449
450 auto inputCpy = mapStateOutputToInput(outputCpy);
451 JLM_ASSERT(inputCpy);
452
453 auto user = users[inputCpy->index()];
454 user->divert_to(&outputCpy);
455 }
456 else
457 {
458 // If we cannot map the output state to the input state of the node, we fall back value-edge
459 // semantic for hoisting.
460 auto & newOutputOrg = rvsdg::RouteToRegion(outputCpy, *node.region());
461 outputOrg.divert_users(&newOutputOrg);
462 }
463 }
464 else if (outputOrg.Type()->Kind() == rvsdg::TypeKind::Value)
465 {
466 auto & newOutputOrg = rvsdg::RouteToRegion(outputCpy, *node.region());
467 outputOrg.divert_users(&newOutputOrg);
468 }
469 else
470 {
471 throw std::logic_error(util::strfmt("Unhandled type kind!"));
472 }
473 }
474}
475
476void
477NodeHoisting::hoistNodes(rvsdg::Region & region)
478{
479 // FIXME: We a routing unnecessary values through gamma and theta nodes. We should cluster
480 // subgraphs that need to be hoisted to avoid unnecessary routing.
481 for (const auto node : rvsdg::TopDownTraverser(&region))
482 {
483 auto & targetRegion = context_->getTargetRegion(*node);
484 if (&targetRegion != node->region())
485 {
486 copyNodeToTargetRegion(*node);
487 }
488
489 // Handle innermost regions
490 if (const auto structuralNode = dynamic_cast<rvsdg::StructuralNode *>(node))
491 {
492 for (auto & subregion : structuralNode->Subregions())
493 {
494 hoistNodes(subregion);
495 }
496 }
497 }
498
499 region.prune(false);
500}
501
502void
503NodeHoisting::printHoistChain(const rvsdg::Region & region) const
504{
505 for (const auto node : rvsdg::TopDownConstTraverser(&region))
506 {
508 *node,
509 [&](const rvsdg::StructuralNode & structuralNode)
510 {
511 for (auto & subregion : structuralNode.Subregions())
512 {
513 printHoistChain(subregion);
514 }
515 },
516 [&](const rvsdg::SimpleNode & simpleNode)
517 {
518 auto & targetRegion = context_->getTargetRegion(simpleNode);
519
520 if (&targetRegion != node->region())
521 {
522 std::cerr << node->DebugString() << "[" << node->GetNodeId() << ", "
523 << node->region()->getRegionId() << "]: ";
524 auto currentRegion = node->region();
525 do
526 {
527 std::cerr << currentRegion->node()->DebugString() << "["
528 << currentRegion->getRegionId() << "] -> ";
529
530 currentRegion = currentRegion->node()->region();
531 } while (currentRegion != &targetRegion);
532
533 std::cerr << currentRegion->node()->DebugString() << "[" << currentRegion->getRegionId()
534 << "]" << std::endl;
535 }
536 },
537 []()
538 {
539 throw std::logic_error("Unhandled node type!");
540 });
541 }
542}
543
544void
545NodeHoisting::hoistNodesInLambda(rvsdg::LambdaNode & lambdaNode)
546{
547 context_ = Context::create(lambdaNode);
548
549 markNodes(*lambdaNode.subregion());
550 // printHoistChain(*lambdaNode.subregion());
551 hoistNodes(*lambdaNode.subregion());
552
553 context_.reset();
554}
555
556void
557NodeHoisting::hoistNodesInRootRegion(rvsdg::Region & region)
558{
559 for (auto & node : rvsdg::TopDownTraverser(&region))
560 {
562 *node,
563 [&](rvsdg::LambdaNode & lambdaNode)
564 {
565 hoistNodesInLambda(lambdaNode);
566 },
567 [&](rvsdg::PhiNode & phiNode)
568 {
569 hoistNodesInRootRegion(*phiNode.subregion());
570 },
571 [](rvsdg::DeltaNode &)
572 {
573 // Nothing needs to be done
574 },
576 {
577 // Nothing needs to be done
578 },
579 [&]()
580 {
581 throw std::logic_error(util::strfmt("Unhandled node type: ", node->DebugString()));
582 });
583 }
584}
585
586void
588{
589 auto statistics = Statistics::Create(rvsdgModule.SourceFilePath().value());
590
591 statistics->start(rvsdgModule.Rvsdg());
592 hoistNodesInRootRegion(rvsdgModule.Rvsdg().GetRootRegion());
593 statistics->end(rvsdgModule.Rvsdg());
594
595 statisticsCollector.CollectDemandedStatistics(std::move(statistics));
596}
597}
static jlm::util::StatisticsCollector statisticsCollector
Definition PullTests.cpp:17
Call operation class.
Definition call.hpp:251
static rvsdg::Input & MapMemoryStateOutputToInput(const rvsdg::Output &output)
Definition Load.hpp:157
Context(rvsdg::LambdaNode &lambdaNode)
void addTargetRegion(const rvsdg::Node &node, rvsdg::Region &region) noexcept
std::unordered_map< const rvsdg::Node *, rvsdg::Region * > TargetRegion_
rvsdg::Region & getTargetRegion(const rvsdg::Node &node) const noexcept
static std::unique_ptr< Context > create(rvsdg::LambdaNode &lambdaNode)
rvsdg::Region & getLambdaSubregion() const noexcept
void end(const rvsdg::Graph &graph) noexcept
Statistics(const util::FilePath &sourceFile)
static std::unique_ptr< Statistics > Create(const util::FilePath &sourceFile)
void start(const rvsdg::Graph &graph) noexcept
Node Hoisting Transformation.
~NodeHoisting() noexcept override
static std::vector< rvsdg::Input * > getUsersFromTargetRegion(rvsdg::Node &node, rvsdg::Region &targetRegion)
rvsdg::Region & computeTargetRegion(const rvsdg::Node &node) const
std::unique_ptr< Context > context_
static bool isInvariantMemoryStateLoopVar(const rvsdg::ThetaNode::LoopVar &loopVar)
void markNodes(const rvsdg::Region &region)
static rvsdg::Input & getUserFromTargetRegion(rvsdg::Input &input, rvsdg::Region &targetRegion)
Conditional operator / pattern matching.
Definition gamma.hpp:99
Region & GetRootRegion() const noexcept
Definition graph.hpp:99
Output * origin() const noexcept
Definition node.hpp:58
Region * region() const noexcept
Definition node.cpp:83
rvsdg::Region * subregion() const noexcept
Definition lambda.hpp:138
NodeOutput * output(size_t index) const noexcept
Definition node.hpp:650
OutputIteratorRange Outputs() noexcept
Definition node.hpp:657
virtual const Operation & GetOperation() const noexcept=0
rvsdg::Region * region() const noexcept
Definition node.hpp:761
InputIteratorRange Inputs() noexcept
Definition node.hpp:622
size_t ninputs() const noexcept
Definition node.hpp:609
size_t noutputs() const noexcept
Definition node.hpp:644
virtual Node * copy(rvsdg::Region *region, const std::vector< jlm::rvsdg::Output * > &operands) const
Definition node.cpp:369
rvsdg::Region * region() const noexcept
Definition node.cpp:151
UsersRange Users()
Definition node.hpp:354
const std::shared_ptr< const rvsdg::Type > & Type() const noexcept
Definition node.hpp:366
size_t nusers() const noexcept
Definition node.hpp:280
A phi node represents the fixpoint of mutually recursive definitions.
Definition Phi.hpp:46
Represent acyclic RVSDG subgraphs.
Definition region.hpp:213
void prune(bool recursive)
Definition region.cpp:326
rvsdg::StructuralNode * node() const noexcept
Definition region.hpp:301
size_t getDepth() const noexcept
Definition region.hpp:291
const std::optional< util::FilePath > & SourceFilePath() const noexcept
Graph & Rvsdg() noexcept
SubregionIteratorRange Subregions()
void CollectDemandedStatistics(std::unique_ptr< Statistics > statistics)
Statistics Interface.
util::Timer & GetTimer(const std::string &name)
util::Timer & AddTimer(std::string name)
void AddMeasurement(std::string name, T value)
void start() noexcept
Definition time.hpp:54
void stop() noexcept
Definition time.hpp:67
#define JLM_ASSERT(x)
Definition common.hpp:16
Global memory state passed between functions.
static bool hasOnlyValueInputs(const rvsdg::Node &node)
static rvsdg::Region & limitTargetRegion(const rvsdg::Node &node, rvsdg::Region &targetRegion)
static rvsdg::Input * mapStateOutputToInput(rvsdg::Output &output)
void MatchTypeWithDefault(T &obj, const Fns &... fns)
Pattern match over subclass type of given object with default handler.
static bool ThetaLoopVarIsInvariant(const ThetaNode::LoopVar &loopVar) noexcept
Definition theta.hpp:266
Output & RouteToRegion(Output &output, Region &region)
Definition node.cpp:381
@ State
Designate a state type.
@ Value
Designate a value type.
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
Definition region.cpp:861
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.
Definition node.hpp:872
detail::TopDownTraverserGeneric< false > TopDownTraverser
Traverser for visiting every node in a region in a top down order.
static std::string strfmt(Args... args)
Definition strfmt.hpp:35
Description of a loop-carried variable.
Definition theta.hpp:50
rvsdg::Output * pre
Variable before iteration (input argument to subregion).
Definition theta.hpp:58
rvsdg::Output * output
Variable at loop exit (output of theta).
Definition theta.hpp:66
rvsdg::Input * post
Variable after iteration (output result from subregion).
Definition theta.hpp:62