Jlm
Loading...
Searching...
No Matches
rhls-dne.cpp
Go to the documentation of this file.
1/*
2 * Copyright 2021 David Metz <david.c.metz@ntnu.no>
3 * See COPYING for terms of redistribution.
4 */
5
6#include <algorithm>
9#include <jlm/hls/ir/hls.hpp>
13
14namespace jlm::hls
15{
16
17static bool
19{
20 bool any_changed = false;
21
22 // Check if any output directly maps to an input into the loop.
23 // If yes, divert users of this output to the pre-loop value.
24 // As a consequence, this output is now unused and will be
25 // removed by remove_unused_loop_outputs (and then lead to
26 // the input also removed by remove_unused_loop_inputs).
27 for (auto exitvar : ln->getExitVars())
28 {
29 auto loopval = exitvar.inner->origin();
31 {
32 auto loopvar = ln->mapArgument(*loopval);
33 if (auto entry = std::get_if<LoopNode::EntryVar>(&loopvar))
34 {
35 exitvar.output->divert_users(entry->input->origin());
36 any_changed = true;
37 }
38 }
39 }
40
41 return any_changed;
42}
43
44static bool
46{
47 // Keep only those entry vars that are not dead.
48 std::vector<LoopNode::ExitVar> vars = ln->getExitVars();
49 vars.erase(
50 std::remove_if(
51 vars.begin(),
52 vars.end(),
53 [](const LoopNode::ExitVar & var)
54 {
55 return !var.output->IsDead();
56 }),
57 vars.end());
58
59 // Remove all dead vars.
60 bool any_changed = !vars.empty();
61 ln->removeExitVars(std::move(vars));
62 return any_changed;
63}
64
65static bool
67{
68 // Keep only those entry vars that are not dead.
69 std::vector<LoopNode::EntryVar> vars = ln->getEntryVars();
70 vars.erase(
71 std::remove_if(
72 vars.begin(),
73 vars.end(),
74 [](const LoopNode::EntryVar & var)
75 {
76 return !var.inner->IsDead();
77 }),
78 vars.end());
79
80 // Remove all dead vars.
81 bool any_changed = !vars.empty();
82 ln->removeEntryVars(std::move(vars));
83 return any_changed;
84}
85
86static bool
88{
89 // Keep only back edge vars that have a user (instead of
90 // simply forwarding to itself).
91 std::vector<LoopNode::BackEdgeVar> vars = ln->getBackEdgeVars();
92 vars.erase(
93 std::remove_if(
94 vars.begin(),
95 vars.end(),
96 [](const LoopNode::BackEdgeVar & var)
97 {
98 return !(var.pre->nusers() == 1 && var.post->origin() == var.pre);
99 }),
100 vars.end());
101 // Remove all that have exactly one user, namely forward itself
102 // to next loop iteration.
103 bool any_changed = !vars.empty();
104 ln->removeBackEdgeVars(std::move(vars));
105 return any_changed;
106}
107
108static bool
110{
111 const auto mux_op = util::assertedCast<const MuxOperation>(&dmux_node->GetOperation());
112 JLM_ASSERT(mux_op->discarding);
113 // check if all inputs have the same origin
114 bool all_inputs_same = true;
115 auto first_origin = dmux_node->input(1)->origin();
116 for (size_t i = 2; i < dmux_node->ninputs(); ++i)
117 {
118 if (dmux_node->input(i)->origin() != first_origin)
119 {
120 all_inputs_same = false;
121 break;
122 }
123 }
124 if (all_inputs_same)
125 {
126 dmux_node->output(0)->divert_users(first_origin);
128 return true;
129 }
130 return false;
131}
132
133static bool
135{
136 auto mux_op = util::assertedCast<const MuxOperation>(&ndmux_node->GetOperation());
137 JLM_ASSERT(!mux_op->discarding);
138 // check if all inputs go to outputs of same branch
139 bool all_inputs_same_branch = true;
140 rvsdg::Node * origin_branch = nullptr;
141 for (size_t i = 1; i < ndmux_node->ninputs(); ++i)
142 {
143 if (auto node = rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(*ndmux_node->input(i)->origin()))
144 {
145 if (dynamic_cast<const BranchOperation *>(&node->GetOperation())
146 && ndmux_node->input(i)->origin()->nusers() == 1)
147 {
148 if (i == 1)
149 {
150 origin_branch = node;
151 continue;
152 }
153 else if (origin_branch == node)
154 {
155 continue;
156 }
157 }
158 }
160 break;
161 }
162 if (all_inputs_same_branch && origin_branch->input(0)->origin() == ndmux_node->input(0)->origin())
163 {
164 // same control origin + all inputs to branch
165 ndmux_node->output(0)->divert_users(origin_branch->input(1)->origin());
167 JLM_ASSERT(origin_branch != nullptr);
169 return true;
170 }
171 return false;
172}
173
174static bool
176{
177 const auto mux_op = util::assertedCast<const MuxOperation>(&ndmux_node->GetOperation());
178 JLM_ASSERT(!mux_op->discarding);
179
180 auto arg = ndmux_node->input(2)->origin();
182 if (!loopNode)
183 {
184 return false;
185 }
186 auto var = loopNode->mapArgument(*arg);
187 // origin is a backedege argument
188 auto backedge = std::get_if<LoopNode::BackEdgeVar>(&var);
189 if (!backedge)
190 {
191 return false;
192 }
193 // one branch
194 if (ndmux_node->output(0)->nusers() != 1)
195 {
196 return false;
197 }
198 auto branch_in_node =
199 rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(*ndmux_node->output(0)->Users().begin());
200 if (!branch_in_node || !dynamic_cast<const BranchOperation *>(&branch_in_node->GetOperation()))
201 {
202 return false;
203 }
204 // one buffer
205 if (branch_in_node->output(1)->nusers() != 1)
206 {
207 return false;
208 }
209 auto buf_in_node =
210 rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(*branch_in_node->output(1)->Users().begin());
211 if (!buf_in_node || !dynamic_cast<const BufferOperation *>(&buf_in_node->GetOperation()))
212 {
213 return false;
214 }
215 auto buf_out = buf_in_node->output(0);
216 if (buf_out != backedge->post->origin())
217 {
218 // no connection back up
219 return false;
220 }
221 // depend on same control
222 auto branch_cond_origin = branch_in_node->input(0)->origin();
223 auto pred_buf_out_node =
226 || !dynamic_cast<const PredicateBufferOperation *>(&pred_buf_out_node->GetOperation()))
227 {
228 return false;
229 }
230 auto pred_buf_cond_origin = pred_buf_out_node->input(0)->origin();
231 // TODO: remove this once predicate buffers decouple combinatorial loops
234 || !dynamic_cast<const BufferOperation *>(&extra_buf_out_node->GetOperation()))
235 {
236 return false;
237 }
238 auto extra_buf_cond_origin = extra_buf_out_node->input(0)->origin();
239
241 {
242 auto var = loopNode->mapArgument(*extra_buf_cond_origin);
243 if (auto extra_be = std::get_if<LoopNode::BackEdgeVar>(&var))
244 {
245 extra_buf_cond_origin = extra_be->post->origin();
246 }
247 }
249 {
250 return false;
251 }
252 // divert users
253 branch_in_node->output(0)->divert_users(ndmux_node->input(1)->origin());
254 buf_out->divert_users(backedge->pre);
258 loopNode->removeBackEdgeVars({ *backedge });
259 return true;
260}
261
262static bool
264{
266
267 // one branch
268 if (lcb_node->output(0)->nusers() != 1)
269 {
270 return false;
271 }
274 if (!branchNode || !branchOperation || !branchOperation->loop)
275 {
276 return false;
277 }
278 // no user
279 if (branchNode->output(1)->nusers())
280 {
281 return false;
282 }
283 // depend on same control
284 auto branch_cond_origin = branchNode->input(0)->origin();
285 auto pred_buf_out = dynamic_cast<rvsdg::NodeOutput *>(lcb_node->input(0)->origin());
286 if (!pred_buf_out
287 || !dynamic_cast<const PredicateBufferOperation *>(&pred_buf_out->node()->GetOperation()))
288 {
289 return false;
290 }
292 // TODO: remove this once predicate buffers decouple combinatorial loops
294 if (!extra_buf_out
295 || !dynamic_cast<const BufferOperation *>(&extra_buf_out->node()->GetOperation()))
296 {
297 return false;
298 }
300
302 if (loopNode)
303 {
304 auto var = loopNode->mapArgument(*extra_buf_cond_origin);
305 if (auto pred_be = std::get_if<LoopNode::BackEdgeVar>(&var))
306 {
307 extra_buf_cond_origin = pred_be->post->origin();
308 }
309 }
311 {
312 return false;
313 }
314 // divert users
315 branchNode->output(0)->divert_users(lcb_node->input(1)->origin());
318 return true;
319}
320
321static bool
323{
324 if (split_node->noutputs() == 1)
325 {
326 split_node->output(0)->divert_users(split_node->input(0)->origin());
327 JLM_ASSERT(split_node->IsDead());
329 return true;
330 }
331 // this merges downward and removes unused outputs (should only exist as a result of eliminating
332 // merges)
333 std::vector<rvsdg::Output *> combined_outputs;
334 for (size_t i = 0; i < split_node->noutputs(); ++i)
335 {
336 if (split_node->output(i)->IsDead())
337 continue;
338 auto user = get_mem_state_user(split_node->output(i));
340 {
342 for (size_t j = 0; j < sub_split->noutputs(); ++j)
343 {
344 combined_outputs.push_back(sub_split->output(j));
345 }
346 }
347 else
348 {
349 combined_outputs.push_back(split_node->output(i));
350 }
351 }
352 if (combined_outputs.size() != split_node->noutputs())
353 {
355 *split_node->input(0)->origin(),
356 combined_outputs.size());
357 for (size_t i = 0; i < combined_outputs.size(); ++i)
358 {
359 combined_outputs[i]->divert_users(new_outputs[i]);
360 }
361 return true;
362 }
363 return false;
364}
365
366static bool
368{
369 // remove single merge
370 if (merge_node->ninputs() == 1)
371 {
372 merge_node->output(0)->divert_users(merge_node->input(0)->origin());
373 JLM_ASSERT(merge_node->IsDead());
375 return true;
376 }
377 std::vector<rvsdg::Output *> combined_origins;
378 std::unordered_set<rvsdg::SimpleNode *> splits;
379 for (size_t i = 0; i < merge_node->ninputs(); ++i)
380 {
381 auto origin = merge_node->input(i)->origin();
383 {
385 for (size_t j = 0; j < sub_merge->ninputs(); ++j)
386 {
387 combined_origins.push_back(sub_merge->input(j)->origin());
388 }
389 }
391 {
392 // ensure that there is only one direct connection to a split.
393 // We need to keep one, so that the optimizations for decouple edges work
394 auto split = rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(*origin);
395 if (!splits.count(split))
396 {
397 splits.insert(split);
398 combined_origins.push_back(origin);
399 }
400 }
401 else
402 {
403 combined_origins.push_back(merge_node->input(i)->origin());
404 }
405 }
406 if (combined_origins.empty())
407 {
408 // if none of the inputs are real keep the first one
409 combined_origins.push_back(merge_node->input(0)->origin());
410 }
411 if (combined_origins.size() != merge_node->ninputs())
412 {
414 merge_node->output(0)->divert_users(new_output);
415 JLM_ASSERT(merge_node->IsDead());
416 return true;
417 }
418 return false;
419}
420
421bool
423 rvsdg::Region & region,
425{
426 bool any_changed = false;
427 bool changed = false;
428 do
429 {
430 changed = false;
431 for (auto & node : rvsdg::BottomUpTraverser(&region))
432 {
433 if (node->IsDead())
434 {
436 {
437 // TODO: fix this once memory connections are explicit
438 continue;
439 }
441 {
442 continue;
443 }
445 {
446 // TODO: fix - this scenario has only stores and should just be optimized away completely
447 continue;
448 }
449 remove(node);
450 changed = true;
451 }
452 else if (dynamic_cast<rvsdg::LambdaNode *>(node))
453 {
454 JLM_UNREACHABLE("This function works on lambda subregions");
455 }
456 else if (auto ln = dynamic_cast<LoopNode *>(node))
457 {
462 changed |= Run(*ln->subregion(), statisticsCollector);
463 }
464 else if (const auto mux = dynamic_cast<const MuxOperation *>(&node->GetOperation()))
465 {
466 if (mux->discarding)
467 {
468 changed |= dead_spec_gamma(node);
469 }
470 else
471 {
472 changed |= dead_nonspec_gamma(node) || dead_loop(node);
473 }
474 }
476 {
477 changed |= dead_loop_lcb(node);
478 }
479 else if (dynamic_cast<const llvm::MemoryStateSplitOperation *>(&node->GetOperation()))
480 {
481 if (fix_mem_split(node))
482 {
483 changed = true;
484 }
485 }
486 else if (dynamic_cast<const llvm::MemoryStateMergeOperation *>(&node->GetOperation()))
487 {
488 if (fix_mem_merge(node))
489 {
490 changed = true;
491 }
492 }
493 if (changed)
494 {
495 // Changes might break bottom up traversal
496 break;
497 }
498 }
500 } while (changed);
501
502 return any_changed;
503}
504
506
510
511void
515{
516 auto & graph = rvsdgModule.Rvsdg();
517 const auto rootRegion = &graph.GetRootRegion();
518 if (rootRegion->numNodes() != 1)
519 {
520 throw util::Error("Root should have only one node now");
521 }
522 const auto lambdaNode =
523 dynamic_cast<const rvsdg::LambdaNode *>(rootRegion->Nodes().begin().ptr());
524 if (!lambdaNode)
525 {
526 throw util::Error("Node needs to be a lambda");
527 }
528 Run(*lambdaNode->subregion(), statisticsCollector);
529}
530
531} // namespace jlm::hls
~RhlsDeadNodeElimination() noexcept override
void Run(rvsdg::RvsdgModule &rvsdgModule, util::StatisticsCollector &statisticsCollector) override
Perform RVSDG transformation.
Definition rhls-dne.cpp:512
static rvsdg::Output * Create(const std::vector< rvsdg::Output * > &operands)
static std::vector< rvsdg::Output * > Create(rvsdg::Output &operand, const size_t numResults)
Output * origin() const noexcept
Definition node.hpp:58
Node * node() const noexcept
Definition node.hpp:572
NodeInput * input(size_t index) const noexcept
Definition node.hpp:615
NodeOutput * output(size_t index) const noexcept
Definition node.hpp:650
void divert_users(jlm::rvsdg::Output *new_origin)
Definition node.hpp:301
Represent acyclic RVSDG subgraphs.
Definition region.hpp:213
Represents an RVSDG transformation.
#define JLM_ASSERT(x)
Definition common.hpp:16
#define JLM_UNREACHABLE(msg)
Definition common.hpp:43
static bool remove_loop_passthrough(LoopNode *ln)
Definition rhls-dne.cpp:18
static bool dead_loop_lcb(rvsdg::Node *lcb_node)
Definition rhls-dne.cpp:263
static bool dead_loop(rvsdg::Node *ndmux_node)
Definition rhls-dne.cpp:175
static bool remove_unused_loop_outputs(LoopNode *ln)
Definition rhls-dne.cpp:45
static bool fix_mem_split(rvsdg::Node *split_node)
Definition rhls-dne.cpp:322
static bool dead_spec_gamma(rvsdg::Node *dmux_node)
Definition rhls-dne.cpp:109
static bool remove_unused_loop_inputs(LoopNode *ln)
Definition rhls-dne.cpp:66
static bool fix_mem_merge(rvsdg::Node *merge_node)
Definition rhls-dne.cpp:367
rvsdg::Input * get_mem_state_user(rvsdg::Output *state_edge)
static bool dead_nonspec_gamma(rvsdg::Node *ndmux_node)
Definition rhls-dne.cpp:134
static bool remove_unused_loop_backedges(LoopNode *ln)
Definition rhls-dne.cpp:87
static util::StatisticsCollector statisticsCollector
static void remove(Node *node)
Definition region.hpp:1035
detail::BottomUpTraverserGeneric< false > BottomUpTraverser
Traverser for visiting every node in a region in a bottom up order.
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.
Definition node.hpp:872
Variable passed between hls loop iterations.
Definition hls.hpp:752
Variable entering the hls loop.
Definition hls.hpp:722
Variable exiting the hls loop.
Definition hls.hpp:737