Jlm
Loading...
Searching...
No Matches
mem-queue.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
10#include <jlm/hls/ir/hls.hpp>
17#include <jlm/rvsdg/gamma.hpp>
20#include <jlm/rvsdg/node.hpp>
21#include <jlm/rvsdg/theta.hpp>
22#include <jlm/rvsdg/view.hpp>
23
24#include <deque>
25
26namespace jlm::hls
27{
28
29static void
32 std::vector<jlm::rvsdg::SimpleNode *> & load_nodes,
33 std::vector<jlm::rvsdg::SimpleNode *> & store_nodes,
34 std::unordered_set<jlm::rvsdg::Output *> & visited)
35{
37 {
38 return;
39 }
40 if (visited.count(op))
41 {
42 // skip already processed outputs
43 return;
44 }
45 visited.insert(op);
46 for (auto & user : op->Users())
47 {
49 user.GetOwner(),
50 [&](rvsdg::Node * node)
51 {
52 rvsdg::MatchTypeOrFail(
53 *node,
54 [&](rvsdg::SimpleNode & simplenode)
55 {
56 rvsdg::MatchType(
57 simplenode.GetOperation(),
58 [&](const jlm::llvm::StoreNonVolatileOperation &)
59 {
60 store_nodes.push_back(&simplenode);
61 },
62 [&](const jlm::llvm::LoadNonVolatileOperation &)
63 {
64 load_nodes.push_back(&simplenode);
65 });
66 for (auto & output : simplenode.Outputs())
67 {
68 find_load_store(&output, load_nodes, store_nodes, visited);
69 }
70 },
71 [&](LoopNode & loop)
72 {
73 find_load_store(loop.mapInput(user).inner, load_nodes, store_nodes, visited);
74 },
75 [&](rvsdg::ThetaNode & theta)
76 {
77 find_load_store(theta.MapInputLoopVar(user).pre, load_nodes, store_nodes, visited);
78 },
79 [&](rvsdg::GammaNode & gamma)
80 {
81 rvsdg::MatchVariant(
82 gamma.MapInput(user),
83 [&](const rvsdg::GammaNode::MatchVar &)
84 {
85 },
87 {
88 for (auto out : evar.branchArgument)
89 {
90 find_load_store(out, load_nodes, store_nodes, visited);
91 }
92 });
93 });
94 },
95 [&](rvsdg::Region * region)
96 {
97 rvsdg::MatchTypeOrFail(
98 *region->node(),
99 [&](LoopNode & loop)
100 {
101 rvsdg::MatchVariant(
102 loop.mapResult(user),
103 [&](const LoopNode::BackEdgeVar & backedge)
104 {
105 find_load_store(backedge.pre, load_nodes, store_nodes, visited);
106 },
107 [&](const LoopNode::ExitVar & exit)
108 {
109 find_load_store(exit.output, load_nodes, store_nodes, visited);
110 });
111 },
112 [&](rvsdg::ThetaNode & theta)
113 {
114 rvsdg::MatchVariant(
115 theta.mapResult(user),
116 [&](const rvsdg::ThetaNode::LoopVar & loopvar)
117 {
118 find_load_store(loopvar.pre, load_nodes, store_nodes, visited);
119 find_load_store(loopvar.output, load_nodes, store_nodes, visited);
120 },
121 [&](const rvsdg::ThetaNode::PredicateVar &)
122 {
123 });
124 },
125 [&](rvsdg::GammaNode & gamma)
126 {
128 gamma.MapBranchResultExitVar(user).output,
129 load_nodes,
130 store_nodes,
131 visited);
132 });
133 });
134 }
135}
136
137static rvsdg::StructuralOutput *
139{
140 auto sti_arg = sti->arguments.first();
141 JLM_ASSERT(sti_arg->nusers() == 1);
142 auto & user = *sti_arg->Users().begin();
143 auto [muxNode, muxOperation] =
145 JLM_ASSERT(muxNode && muxOperation);
146 for (size_t i = 1; i < 3; ++i)
147 {
148 auto arg = muxNode->input(i)->origin();
149 auto loopNode = rvsdg::TryGetRegionParentNode<LoopNode>(*arg);
150 if (!loopNode)
151 {
152 continue;
153 }
154 auto var = loopNode->mapArgument(*arg);
155 if (auto ba = std::get_if<LoopNode::BackEdgeVar>(&var))
156 {
157 auto res = ba->post;
158 JLM_ASSERT(res);
159 auto [bufferNode, bufferOperation] =
161 JLM_ASSERT(bufferNode && bufferOperation);
162 auto [branchNode, branchOperation] =
164 *bufferNode->input(0)->origin());
165 JLM_ASSERT(branchNode && branchOperation);
166 for (size_t j = 0; j < 2; ++j)
167 {
168 JLM_ASSERT(branchNode->output(j)->nusers() == 1);
169 auto result =
170 dynamic_cast<jlm::rvsdg::RegionResult *>(&*branchNode->output(j)->Users().begin());
171 if (result)
172 {
173 return result->output();
174 }
175 }
176 }
177 }
178 JLM_UNREACHABLE("This should never happen");
179}
180
181static rvsdg::Output *
183 jlm::rvsdg::Output * mem_edge,
184 jlm::rvsdg::Output * addr_edge,
186 jlm::rvsdg::Output ** new_mem_edge,
187 std::vector<jlm::rvsdg::Output *> & store_addresses,
188 std::vector<jlm::rvsdg::Output *> & store_dequeues,
189 std::vector<bool> & store_precedes,
190 bool * load_encountered)
191{
192 // follows along mem edge and routes addr edge through the same regions
193 // redirects the supplied load to the new edge and adds it to stores
194 // the new edge might be routed through unnecessary regions. This should be fixed by running DNE
195 while (true)
196 {
197 // each iteration should update common_edge and/or new_edge
198 JLM_ASSERT(mem_edge->nusers() == 1);
199 JLM_ASSERT(addr_edge->nusers() == 1);
200 JLM_ASSERT(mem_edge != addr_edge);
201 JLM_ASSERT(mem_edge->region() == addr_edge->region());
202 auto user = &*mem_edge->Users().begin();
203 auto & addr_edge_user = *addr_edge->Users().begin();
204 if (dynamic_cast<jlm::rvsdg::RegionResult *>(user))
205 {
206 JLM_UNREACHABLE("THIS SHOULD NOT HAPPEN");
207 // end of region reached
208 }
209 else if (auto sti = dynamic_cast<jlm::rvsdg::StructuralInput *>(user))
210 {
211 auto loop_node = jlm::util::assertedCast<jlm::hls::LoopNode>(sti->node());
212 jlm::rvsdg::Output * buffer = nullptr;
213 auto addr_edge_before_loop = addr_edge;
214 addr_edge = loop_node->AddLoopVar(addr_edge, &buffer);
215 addr_edge_user.divert_to(addr_edge);
216 mem_edge = find_loop_output(sti);
217 auto sti_arg = sti->arguments.first();
218 JLM_ASSERT(sti_arg->nusers() == 1);
219 auto & user = *sti_arg->Users().begin();
220 auto [muxNode, muxOperation] =
222 JLM_ASSERT(muxNode && muxOperation);
223 JLM_ASSERT(buffer->nusers() == 1);
224 // use a separate vector to check if the loop contains stores
225 std::vector<jlm::rvsdg::Output *> loop_store_addresses;
227 muxNode->output(0),
228 buffer,
229 load,
230 nullptr,
231 loop_store_addresses,
232 store_dequeues,
233 store_precedes,
234 load_encountered);
235 if (loop_store_addresses.empty())
236 {
237 jlm::hls::convert_loop_state_to_lcb(&*addr_edge_before_loop->Users().begin());
238 }
239 else
240 {
241 store_addresses.insert(
242 store_addresses.cend(),
243 loop_store_addresses.begin(),
244 loop_store_addresses.end());
245 }
246 }
248 {
249 auto op = &sn->GetOperation();
250
251 if (auto br = dynamic_cast<const jlm::hls::BranchOperation *>(op))
252 {
253 if (!br->loop)
254 {
255 // start of gamma
256 auto load_branch_out =
257 jlm::hls::BranchOperation::create(*sn->input(0)->origin(), *addr_edge, false);
258 for (size_t i = 0; i < sn->noutputs(); ++i)
259 {
260 // dummy user for edge
261 auto dummy_user_tmp = jlm::hls::SinkOperation::create(*load_branch_out[i]);
262 // Sink ops doesn't have any outputs so we get an empty vector back
263 // But we are not allowed to discard the vector and can't have unused variables
264 // So adding a meaningless assert to get it to compile
265 JLM_ASSERT(dummy_user_tmp.size() == 0);
267 *load_branch_out[i]->Users().begin());
268 // need both load and common edge here
269 load_branch_out[i] = separate_load_edge(
270 sn->output(i),
271 load_branch_out[i],
272 load,
273 &mem_edge,
274 store_addresses,
275 store_dequeues,
276 store_precedes,
277 load_encountered);
278 JLM_ASSERT(load_branch_out[i]->nusers() == 1);
279 JLM_ASSERT(dummy_user->input(0)->origin() == load_branch_out[i]);
280 remove(dummy_user);
281 }
282 // create mux
283 JLM_ASSERT(mem_edge->nusers() == 1);
284 auto [muxNode, muxOperation] =
286 *mem_edge->Users().begin());
287 JLM_ASSERT(muxNode && muxOperation);
289 *muxNode->input(0)->origin(),
290 load_branch_out,
291 muxOperation->discarding,
292 false)[0];
293 addr_edge_user.divert_to(addr_edge);
294 mem_edge = muxNode->output(0);
295 }
296 else
297 {
298 // end of loop
300 return nullptr;
301 }
302 }
303 else if (auto mx = dynamic_cast<const jlm::hls::MuxOperation *>(op))
304 {
305 JLM_ASSERT(!mx->loop);
306 // end of gamma
307 JLM_ASSERT(new_mem_edge);
308 *new_mem_edge = mem_edge;
309 return addr_edge;
310 }
311 else if (dynamic_cast<const jlm::llvm::StoreNonVolatileOperation *>(op))
312 {
313 auto sg_out = jlm::hls::StateGateOperation::create(*sn->input(0)->origin(), { addr_edge });
314 addr_edge = sg_out[1];
315 addr_edge_user.divert_to(addr_edge);
316 store_addresses.push_back(jlm::hls::route_to_region_rhls((*load)->region(), sg_out[0]));
317 store_precedes.push_back(!*load_encountered);
318 mem_edge = sn->output(0);
319 JLM_ASSERT(mem_edge->nusers() == 1);
320 user = &*mem_edge->Users().begin();
321 auto [mssNode, msso] =
323 if (mssNode && msso)
324 {
325 // handle case where output of store is already connected to a MemStateSplit by adding an
326 // output
327 auto store_split =
328 jlm::llvm::MemoryStateSplitOperation::Create(*mem_edge, msso->nresults() + 1);
329 for (size_t i = 0; i < msso->nresults(); ++i)
330 {
331 mssNode->output(i)->divert_users(store_split[i]);
332 }
333 remove(mssNode);
334 mem_edge = store_split[0];
335 store_dequeues.push_back(
336 jlm::hls::route_to_region_rhls((*load)->region(), store_split.back()));
337 }
338 else
339 {
340 auto store_split = jlm::llvm::MemoryStateSplitOperation::Create(*mem_edge, 2);
341 mem_edge = store_split[0];
342 user->divert_to(mem_edge);
343 store_dequeues.push_back(
344 jlm::hls::route_to_region_rhls((*load)->region(), store_split[1]));
345 }
346 }
347 else if (auto lo = dynamic_cast<const jlm::llvm::LoadNonVolatileOperation *>(op))
348 {
349 JLM_ASSERT(sn->noutputs() == 2);
350 if (sn == *load)
351 {
352 // create state gate for addr edge
353 auto addr_sg_out =
354 jlm::hls::StateGateOperation::create(*sn->input(0)->origin(), { addr_edge });
355 addr_edge = addr_sg_out[1];
356 addr_edge_user.divert_to(addr_edge);
357 auto addr_sg_out2 = jlm::hls::StateGateOperation::create(*addr_sg_out[0], { addr_edge });
358 addr_edge = addr_sg_out2[1];
359 addr_edge_user.divert_to(addr_edge);
360 // remove state edges from load
361 auto new_load_outputs = jlm::llvm::LoadNonVolatileOperation::Create(
362 addr_sg_out2[0],
363 {},
364 lo->GetLoadedType(),
365 lo->GetAlignment());
366 // create state gate for mem edge and load data
367 auto mem_sg_out =
368 jlm::hls::StateGateOperation::create(*new_load_outputs[0], { mem_edge });
369 mem_edge = mem_sg_out[1];
370
371 sn->output(0)->divert_users(new_load_outputs[0]);
372 user->divert_to(addr_edge);
373 sn->output(1)->divert_users(mem_edge);
374 remove(sn);
375 *load = &jlm::rvsdg::AssertGetOwnerNode<jlm::rvsdg::SimpleNode>(*new_load_outputs[0]);
376 *load_encountered = true;
377 }
378 else
379 {
380 mem_edge = sn->output(1);
381 }
382 }
383 else if (dynamic_cast<const jlm::hls::StateGateOperation *>(op))
384 {
385 mem_edge = sn->output(1);
386 }
387 else if (dynamic_cast<const jlm::llvm::CallOperation *>(op))
388 {
389 JLM_ASSERT("Decoupled nodes not implemented yet");
390 }
391 else if (dynamic_cast<const jlm::llvm::MemoryStateMergeOperation *>(op))
392 {
393 auto si_load_user = jlm::rvsdg::TryGetOwnerNode<jlm::rvsdg::SimpleNode>(addr_edge_user);
395 if (si_load_user && &userNode == sn)
396 {
397 return nullptr;
398 }
399 // TODO: handle
400 JLM_UNREACHABLE("THIS SHOULD NOT HAPPEN");
401 }
402 else
403 {
404 JLM_UNREACHABLE("THIS SHOULD NOT HAPPEN");
405 }
406 }
407 else
408 {
409 JLM_UNREACHABLE("THIS SHOULD NOT HAPPEN");
410 }
411 }
412}
413
416{
417 while (true)
418 {
419 // each iteration should update state_edge
420 JLM_ASSERT(state_edge->nusers() == 1);
421 auto & user = *state_edge->Users().begin();
422 if (dynamic_cast<jlm::rvsdg::RegionResult *>(&user))
423 {
424 // End of region reached
425 return user.origin();
426 }
428 {
429 auto op = &sn->GetOperation();
430 auto br = dynamic_cast<const jlm::hls::BranchOperation *>(op);
431 if (br && !br->loop)
432 {
433 // start of gamma
434 for (size_t i = 0; i < sn->noutputs(); ++i)
435 {
436 state_edge = process_loops(sn->output(i));
437 }
438 }
440 {
441 // end of gamma
442 JLM_ASSERT(sn->noutputs() == 1);
443 return sn->output(0);
444 }
445 else if (dynamic_cast<const jlm::llvm::LambdaExitMemoryStateMergeOperation *>(op))
446 {
447 // end of lambda
448 JLM_ASSERT(sn->noutputs() == 1);
449 return sn->output(0);
450 }
451 else if (dynamic_cast<const jlm::llvm::LoadNonVolatileOperation *>(op))
452 {
453 // load
454 JLM_ASSERT(sn->noutputs() == 2);
455 state_edge = sn->output(1);
456 }
457 else if (dynamic_cast<const jlm::llvm::CallOperation *>(op))
458 {
459 state_edge = sn->output(sn->noutputs() - 1);
460 }
461 else
462 {
463 JLM_ASSERT(sn->noutputs() == 1);
464 state_edge = sn->output(0);
465 }
466 }
467 else if (auto sti = dynamic_cast<jlm::rvsdg::StructuralInput *>(&user))
468 {
469 JLM_ASSERT(dynamic_cast<const jlm::hls::LoopNode *>(sti->node()));
470 // update to output of loop
471 auto mem_edge_after_loop = find_loop_output(sti);
472 JLM_ASSERT(mem_edge_after_loop->nusers() == 1);
473 auto & common_user = *mem_edge_after_loop->Users().begin();
474
475 std::vector<jlm::rvsdg::SimpleNode *> load_nodes;
476 std::vector<jlm::rvsdg::SimpleNode *> store_nodes;
477 std::unordered_set<jlm::rvsdg::Output *> visited;
478 // this is a hack to keep search within the loop
479 visited.insert(mem_edge_after_loop);
480 find_load_store(&*sti->arguments.begin(), load_nodes, store_nodes, visited);
481 auto split_states =
482 jlm::llvm::MemoryStateSplitOperation::Create(*sti->origin(), load_nodes.size() + 1);
483 // handle common edge
484 auto mem_edge = split_states[0];
485 sti->divert_to(mem_edge);
486 split_states[0] = mem_edge_after_loop;
487 state_edge = jlm::llvm::MemoryStateMergeOperation::Create(split_states);
488 common_user.divert_to(state_edge);
489 for (size_t i = 0; i < load_nodes.size(); ++i)
490 {
491 auto load = load_nodes[i];
492 auto addr_edge = split_states[1 + i];
493 std::vector<jlm::rvsdg::Output *> store_addresses;
494 std::vector<jlm::rvsdg::Output *> store_dequeues;
495 std::vector<bool> store_precedes;
496 bool load_encountered = false;
498 mem_edge,
499 addr_edge,
500 &load,
501 nullptr,
502 store_addresses,
503 store_dequeues,
504 store_precedes,
505 &load_encountered);
506 JLM_ASSERT(load_encountered);
507 JLM_ASSERT(store_nodes.size() == store_addresses.size());
508 JLM_ASSERT(store_nodes.size() == store_dequeues.size());
509 auto state_gate_addr_in =
511 .input(0);
512 for (size_t j = 0; j < store_nodes.size(); ++j)
513 {
514 JLM_ASSERT(state_gate_addr_in->origin()->region() == store_addresses[j]->region());
515 JLM_ASSERT(store_dequeues[j]->region() == store_addresses[j]->region());
516 state_gate_addr_in->divert_to(jlm::hls::AddressQueueOperation::create(
517 *state_gate_addr_in->origin(),
518 *store_addresses[j],
519 *store_dequeues[j],
520 store_precedes[j]));
521 }
522 }
523 }
524 else
525 {
526 JLM_UNREACHABLE("THIS SHOULD NOT HAPPEN");
527 }
528 }
529}
530
531static void
533{
534 const auto & graph = rvsdgModule.Rvsdg();
535 const auto rootRegion = &graph.GetRootRegion();
536 if (rootRegion->numNodes() != 1)
537 {
538 throw std::logic_error("Root should have only one node now");
539 }
540
541 const auto lambda = dynamic_cast<const rvsdg::LambdaNode *>(rootRegion->Nodes().begin().ptr());
542 if (!lambda)
543 {
544 throw std::logic_error("Node needs to be a lambda");
545 }
546
547 auto state_arg = &llvm::GetMemoryStateRegionArgument(*lambda);
548 if (!state_arg)
549 {
550 // No memstate, i.e., no memory used
551 return;
552 }
553 // for each state edge:
554 // for each outer loop (theta/loop in lambda region):
555 // split state edge before the loop
556 // * one edge for only stores (preserves store order)
557 // * a separate edge for each load, going through the stores as well
558 // merge state edges after the loop
559 // for each load:
560 // insert store address queue before address input of load
561 // * enq order of stores guaranteed by load edge, deq by store edge
562 // for each store:
563 // insert state gate addr enq + deq after store complete
564
565 // Check if there exists a memory state splitter
566 if (state_arg->nusers() == 1)
567 {
568 auto entryNode = rvsdg::TryGetOwnerNode<rvsdg::Node>(*state_arg->Users().begin());
570 entryNode->GetOperation()))
571 {
572 for (size_t i = 0; i < entryNode->noutputs(); ++i)
573 {
574 // Process each state edge separately
575 jlm::rvsdg::Output * stateEdge = entryNode->output(i);
576 process_loops(stateEdge);
577 }
578 return;
579 }
580 }
581 // There is no memory state splitter, so process the single state edge in the graph
582 process_loops(state_arg);
583}
584
585AddressQueueInsertion::~AddressQueueInsertion() noexcept = default;
586
590
591void
592AddressQueueInsertion::Run(rvsdg::RvsdgModule & rvsdgModule, util::StatisticsCollector &)
593{
594 mem_queue(rvsdgModule);
595}
596
597}
static jlm::rvsdg::Output * create(jlm::rvsdg::Output &check, jlm::rvsdg::Output &enq, jlm::rvsdg::Output &deq, bool combinatorial, size_t capacity=10)
Definition hls.hpp:1097
static std::vector< jlm::rvsdg::Output * > create(jlm::rvsdg::Output &predicate, jlm::rvsdg::Output &value, bool loop=false)
Definition hls.hpp:68
static std::vector< jlm::rvsdg::Output * > create(jlm::rvsdg::Output &predicate, const std::vector< jlm::rvsdg::Output * > &alternatives, bool discarding, bool loop=false)
Definition hls.hpp:235
static std::vector< jlm::rvsdg::Output * > create(jlm::rvsdg::Output &value)
Definition hls.hpp:302
static std::vector< jlm::rvsdg::Output * > create(jlm::rvsdg::Output &addr, const std::vector< jlm::rvsdg::Output * > &states)
Definition hls.hpp:1157
Call operation class.
Definition call.hpp:251
static std::unique_ptr< llvm::ThreeAddressCode > Create(const Variable *address, const Variable *state, std::shared_ptr< const rvsdg::Type > loadedType, size_t alignment)
Definition Load.hpp:448
static rvsdg::Output * Create(const std::vector< rvsdg::Output * > &operands)
static std::vector< rvsdg::Output * > Create(rvsdg::Output &operand, const size_t numResults)
Conditional operator / pattern matching.
Definition gamma.hpp:99
Region & GetRootRegion() const noexcept
Definition graph.hpp:99
rvsdg::Region * region() const noexcept
Definition node.cpp:151
UsersRange Users()
Definition node.hpp:354
void divert_users(jlm::rvsdg::Output *new_origin)
Definition node.hpp:301
size_t nusers() const noexcept
Definition node.hpp:280
Represents the result of a region.
Definition region.hpp:120
StructuralOutput * output() const noexcept
Definition region.hpp:149
Graph & Rvsdg() noexcept
Represents an RVSDG transformation.
ElementType * first() const noexcept
#define JLM_ASSERT(x)
Definition common.hpp:16
#define JLM_UNREACHABLE(msg)
Definition common.hpp:43
static void find_load_store(jlm::rvsdg::Output *op, std::vector< jlm::rvsdg::SimpleNode * > &load_nodes, std::vector< jlm::rvsdg::SimpleNode * > &store_nodes, std::unordered_set< jlm::rvsdg::Output * > &visited)
Definition mem-queue.cpp:30
rvsdg::Output * route_to_region_rhls(rvsdg::Region *target, rvsdg::Output *out)
jlm::rvsdg::Output * process_loops(jlm::rvsdg::Output *state_edge)
static rvsdg::StructuralOutput * find_loop_output(jlm::rvsdg::StructuralInput *sti)
static rvsdg::Output * separate_load_edge(jlm::rvsdg::Output *mem_edge, jlm::rvsdg::Output *addr_edge, jlm::rvsdg::SimpleNode **load, jlm::rvsdg::Output **new_mem_edge, std::vector< jlm::rvsdg::Output * > &store_addresses, std::vector< jlm::rvsdg::Output * > &store_dequeues, std::vector< bool > &store_precedes, bool *load_encountered)
static void mem_queue(rvsdg::RvsdgModule &rvsdgModule)
void convert_loop_state_to_lcb(rvsdg::Input *loop_state_input)
decltype(auto) MatchVariant(T &&obj, Fns &&... fns)
Pattern match over variant.
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.
Definition node.hpp:872
A variable routed into all gamma regions.
Definition gamma.hpp:131