Jlm
Loading...
Searching...
No Matches
mem-conv.cpp
Go to the documentation of this file.
1/*
2 * Copyright 2021 David Metz <david.c.metz@ntnu.no>
3 * and Magnus Sjalander <work@sjalander.com>
4 * See COPYING for terms of redistribution.
5 */
6
12#include <jlm/hls/ir/hls.hpp>
19#include <jlm/rvsdg/gamma.hpp>
23#include <jlm/rvsdg/theta.hpp>
25#include <jlm/rvsdg/view.hpp>
26
27namespace jlm::hls
28{
29rvsdg::SimpleNode *
31 const rvsdg::LambdaNode * lambda,
33{
34 auto response_functions = find_function_arguments(lambda, "decouple_res");
35 for (auto & func : response_functions)
36 {
37 std::unordered_set<rvsdg::Output *> visited;
38 std::vector<rvsdg::SimpleNode *> reponse_calls;
40 for (auto & rc : reponse_calls)
41 {
42 auto response_constant = trace_constant(rc->input(1)->origin());
44 {
45 return rc;
46 }
47 }
48 }
49 JLM_UNREACHABLE("No response found");
50}
51
52static std::pair<rvsdg::Input *, std::vector<rvsdg::Input *>>
54{
55 std::vector<rvsdg::Input *> encountered_muxes;
56 // should encounter no new loops, or gammas, only exit them
58 while (true)
59 {
60 // make sure we make progress
62 if (dynamic_cast<jlm::rvsdg::RegionResult *>(state_edge))
63 {
64 JLM_UNREACHABLE("this should be handled by branch");
65 }
67 {
68 JLM_UNREACHABLE("there should be no new loops");
69 }
70 auto si = state_edge;
76 {
77 // end of loop
80 util::assertedCast<rvsdg::RegionResult>(get_mem_state_user(sn->output(0)))->output());
81 }
82 else if (muxOperation && !muxOperation->loop)
83 {
84 // end of gamma
85 encountered_muxes.push_back(si);
86 state_edge = get_mem_state_user(sn->output(0));
87 }
88 else if (
91 {
92 return { state_edge, encountered_muxes };
93 }
94 else
95 {
96 JLM_UNREACHABLE("whoops");
97 }
98 }
99}
100
101void
103{
104 // replace other branches with undefs, so the stateedge before the res can be killed.
107 for (auto si : encountered_muxes)
108 {
110 for (size_t i = 1; i < sn.ninputs(); ++i)
111 {
112 if (i != si->index())
113 {
114 auto state_dummy = llvm::UndefValueOperation::Create(*si->region(), si->Type());
115 sn.input(i)->divert_to(state_dummy);
116 }
117 }
118 }
119}
120
121void
123{
124 // there is no reason to wait for requests, if we already wait for responses, so we kill the rest
125 // of this state edge
129 std::vector<rvsdg::Output *> merge_origins;
130 for (size_t i = 0; i < merge_node.ninputs(); ++i)
131 {
132 if (i != merge_in->index())
133 {
134 merge_origins.push_back(merge_node.input(i)->origin());
135 }
136 }
138 merge_node.output(0)->divert_users(new_merge_output);
139 JLM_ASSERT(merge_node.IsDead());
141}
142
145 const rvsdg::LambdaNode * lambda,
148{
149 JLM_ASSERT(dynamic_cast<const llvm::CallOperation *>(&decouple_request->GetOperation()));
150 auto channel = decouple_request->input(1)->origin();
152
154
155 // handle request
156 auto addr = decouple_request->input(2)->origin();
157 auto req_mem_state = decouple_request->input(decouple_request->ninputs() - 1)->origin();
158 // state gate for req
160 addr = sg_out[0];
162 // redirect memstate - iostate output has already been removed by mem-sep pass
163 decouple_request->output(decouple_request->noutputs() - 1)->divert_users(req_mem_state);
164
165 // handle response
166 int load_capacity = 10;
168 {
169 auto constant = trace_constant(decouple_response->input(2)->origin());
170 load_capacity = constant->Representation().to_int();
171 assert(load_capacity >= 0);
172 }
176
178 decouple_response->output(0)->divert_users(routed_data);
179 auto response_state_origin = decouple_response->input(decouple_response->ninputs() - 1)->origin();
180
181 if (decouple_request->region() != decouple_response->region())
182 {
183 // they are in different regions, so we handle state edge at response
185 *response_state_origin->region(),
186 response_state_origin->Type());
188 decouple_response->output(decouple_response->noutputs() - 1)->divert_users(sg_resp[1]);
189 JLM_ASSERT(decouple_response->IsDead());
191 JLM_ASSERT(decouple_request->IsDead());
193
196 }
197 else
198 {
199 // they are in the same region, handle at request
200 // remove mem state from response call
201 decouple_response->output(decouple_response->noutputs() - 1)
202 ->divert_users(response_state_origin);
203
205 *response_state_origin->region(),
206 response_state_origin->Type());
207 // put state gate on load response
208 auto sg_resp = StateGateOperation::create(*dload_node->input(1)->origin(), { state_dummy });
209 dload_node->input(1)->divert_to(sg_resp[0]);
211 state_user->divert_to(sg_resp[1]);
212
213 JLM_ASSERT(decouple_response->IsDead());
215 JLM_ASSERT(decouple_request->IsDead());
217
218 // these are swapped in this scenario, since we keep the one from request
221 }
222
223 auto nn = dynamic_cast<rvsdg::NodeOutput *>(dload_out[0])->node();
224 return dynamic_cast<rvsdg::SimpleNode *>(nn);
225}
226
227void
229 rvsdg::Region * region,
230 std::vector<rvsdg::Node *> & loadNodes,
231 std::vector<rvsdg::Node *> & storeNodes,
232 std::vector<rvsdg::Node *> & decoupleNodes,
233 std::unordered_set<rvsdg::Node *> exclude)
234{
235 for (auto & node : rvsdg::TopDownTraverser(region))
236 {
237 if (auto structnode = dynamic_cast<rvsdg::StructuralNode *>(node))
238 {
239 for (size_t n = 0; n < structnode->nsubregions(); n++)
240 gather_mem_nodes(structnode->subregion(n), loadNodes, storeNodes, decoupleNodes, exclude);
241 }
242 else if (auto simplenode = dynamic_cast<rvsdg::SimpleNode *>(node))
243 {
244 if (exclude.find(simplenode) != exclude.end())
245 {
246 continue;
247 }
248 if (dynamic_cast<const llvm::StoreNonVolatileOperation *>(&simplenode->GetOperation()))
249 {
250 storeNodes.push_back(simplenode);
251 }
252 else if (dynamic_cast<const llvm::LoadNonVolatileOperation *>(&simplenode->GetOperation()))
253 {
254 loadNodes.push_back(simplenode);
255 }
256 else if (dynamic_cast<const llvm::CallOperation *>(&simplenode->GetOperation()))
257 {
258 // we only want to collect requests
260 decoupleNodes.push_back(simplenode);
261 }
262 }
263 }
264}
265
273static void
275 rvsdg::Output * output,
276 std::unordered_set<rvsdg::Output *> & visited,
278{
279 if (!rvsdg::is<llvm::PointerType>(output->Type()))
280 {
281 // Only process pointer outputs
282 return;
283 }
284 if (visited.count(output))
285 {
286 // Skip already processed outputs
287 return;
288 }
289 visited.insert(output);
290 for (auto & user : output->Users())
291 {
293 user.GetOwner(),
294 [&](rvsdg::Node * node)
295 {
296 rvsdg::MatchTypeOrFail(
297 *node,
298 [&](rvsdg::SimpleNode & simplenode)
299 {
300 rvsdg::MatchTypeWithDefault(
301 simplenode.GetOperation(),
302 [&](const llvm::StoreNonVolatileOperation &)
303 {
304 tracedPointerNodes.storeNodes.push_back(&simplenode);
305 },
306 [&](const llvm::LoadNonVolatileOperation &)
307 {
308 tracedPointerNodes.loadNodes.push_back(&simplenode);
309 },
310 [&](const llvm::CallOperation &)
311 {
312 JLM_ASSERT(is_dec_req(&simplenode));
313 tracedPointerNodes.decoupleNodes.push_back(&simplenode);
314 },
315 [&]()
316 {
317 for (size_t i = 0; i < simplenode.noutputs(); ++i)
318 {
319 TracePointer(simplenode.output(i), visited, tracedPointerNodes);
320 }
321 });
322 },
323 [&](LoopNode & loop)
324 {
325 TracePointer(loop.mapInput(user).inner, visited, tracedPointerNodes);
326 },
327 [&](rvsdg::ThetaNode & theta)
328 {
329 TracePointer(theta.MapInputLoopVar(user).pre, visited, tracedPointerNodes);
330 },
331 [&](rvsdg::GammaNode & gamma)
332 {
333 rvsdg::MatchVariant(
334 gamma.MapInput(user),
335 [&](const rvsdg::GammaNode::MatchVar &)
336 {
337 },
339 {
340 for (auto arg : evar.branchArgument)
341 {
342 TracePointer(arg, visited, tracedPointerNodes);
343 }
344 });
345 });
346 },
347 [&](rvsdg::Region * region)
348 {
349 rvsdg::MatchTypeOrFail(
350 *region->node(),
351 [&](LoopNode & loop)
352 {
353 rvsdg::MatchVariant(
354 loop.mapResult(user),
355 [&](const LoopNode::BackEdgeVar & backedge)
356 {
357 TracePointer(backedge.pre, visited, tracedPointerNodes);
358 },
359 [&](const LoopNode::ExitVar & exit)
360 {
361 TracePointer(exit.output, visited, tracedPointerNodes);
362 });
363 },
364 [&](rvsdg::ThetaNode & theta)
365 {
366 rvsdg::MatchVariant(
367 theta.mapResult(user),
368 [&](const rvsdg::ThetaNode::LoopVar & loopvar)
369 {
370 TracePointer(loopvar.output, visited, tracedPointerNodes);
371 },
372 [&](const rvsdg::ThetaNode::PredicateVar &)
373 {
374 });
375 },
376 [&](rvsdg::GammaNode & gamma)
377 {
379 gamma.MapBranchResultExitVar(user).output,
380 visited,
381 tracedPointerNodes);
382 });
383 });
384 }
385}
386
387std::vector<TracedPointerNodes>
389{
390 std::vector<TracedPointerNodes> tracedPointerNodes;
391 for (const auto argument : lambda->GetFunctionArguments())
392 {
393 if (rvsdg::is<llvm::PointerType>(argument->Type()))
394 {
395 std::unordered_set<rvsdg::Output *> visited;
396 tracedPointerNodes.emplace_back();
397 TracePointer(argument, visited, tracedPointerNodes.back());
398 }
399 }
400
401 for (auto cv : lambda->GetContextVars())
402 {
403 if (rvsdg::is<llvm::PointerType>(cv.inner->Type()) && !is_function_argument(cv))
404 {
405 std::unordered_set<rvsdg::Output *> visited;
406 tracedPointerNodes.emplace_back();
407 TracePointer(cv.inner, visited, tracedPointerNodes.back());
408 }
409 }
410
411 return tracedPointerNodes;
412}
413
416{
417 if (auto l = dynamic_cast<rvsdg::LambdaNode *>(region->node()))
418 {
419 return l;
420 }
421 return find_containing_lambda(region->node()->region());
422}
423
424static size_t
425CalculatePortWidth(const TracedPointerNodes & tracedPointerNodes)
426{
427 int max_width = 0;
428 for (auto node : tracedPointerNodes.loadNodes)
429 {
430 auto loadOp = util::assertedCast<const llvm::LoadNonVolatileOperation>(&node->GetOperation());
431 auto sz = JlmSize(loadOp->GetLoadedType().get());
432 max_width = sz > max_width ? sz : max_width;
433 }
434 for (auto node : tracedPointerNodes.storeNodes)
435 {
436 auto storeOp = util::assertedCast<const llvm::StoreNonVolatileOperation>(&node->GetOperation());
437 auto sz = JlmSize(&storeOp->GetStoredType());
438 max_width = sz > max_width ? sz : max_width;
439 }
440 for (auto decoupleRequest : tracedPointerNodes.decoupleNodes)
441 {
442 auto lambda = find_containing_lambda(decoupleRequest->region());
443 auto channel = decoupleRequest->input(1)->origin();
444 auto channelConstant = trace_constant(channel);
445 auto reponse = find_decouple_response(lambda, channelConstant);
446 auto sz = JlmSize(reponse->output(0)->Type().get());
447 max_width = sz > max_width ? sz : max_width;
448 }
449 JLM_ASSERT(max_width != 0);
450 return max_width;
451}
452
453static rvsdg::SimpleNode *
456 const rvsdg::Node * originalLoad,
457 rvsdg::Output * response)
458{
459 // We have the load from the original lambda since it is needed to update the smap
460 // We need the load in the new lambda such that we can replace it with a load node with explicit
461 // memory ports
462 auto replacedLoad =
463 &rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.lookup(*originalLoad->output(0)));
464
465 auto loadAddress = replacedLoad->input(0)->origin();
466 std::vector<rvsdg::Output *> states;
467 for (size_t i = 1; i < replacedLoad->ninputs(); ++i)
468 {
469 states.push_back(replacedLoad->input(i)->origin());
470 }
471
472 rvsdg::Node * newLoad = nullptr;
473 if (states.empty())
474 {
475 size_t load_capacity = 10;
476 auto outputs = DecoupledLoadOperation::create(*loadAddress, *response, load_capacity);
477 newLoad = dynamic_cast<rvsdg::NodeOutput *>(outputs[0])->node();
478 }
479 else
480 {
481 // TODO: switch this to a decoupled load?
482 auto outputs = LoadOperation::create(*loadAddress, states, *response);
483 newLoad = dynamic_cast<rvsdg::NodeOutput *>(outputs[0])->node();
484 }
485
486 for (size_t i = 0; i < replacedLoad->noutputs(); ++i)
487 {
488 smap.insert(originalLoad->output(i), newLoad->output(i));
489 replacedLoad->output(i)->divert_users(newLoad->output(i));
490 }
491 remove(replacedLoad);
492 return dynamic_cast<rvsdg::SimpleNode *>(newLoad);
493}
494
495static rvsdg::SimpleNode *
498 const rvsdg::Node * originalStore,
499 rvsdg::Output * response)
500{
501 // We have the store from the original lambda since it is needed to update the smap
502 // We need the store in the new lambda such that we can replace it with a store node with explicit
503 // memory ports
504 auto replacedStore =
505 &rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.lookup(*originalStore->output(0)));
506
507 auto addr = replacedStore->input(0)->origin();
508 JLM_ASSERT(rvsdg::is<llvm::PointerType>(addr->Type()));
509 auto data = replacedStore->input(1)->origin();
510 std::vector<rvsdg::Output *> states;
511 for (size_t i = 2; i < replacedStore->ninputs(); ++i)
512 {
513 states.push_back(replacedStore->input(i)->origin());
514 }
515 auto storeOuts = StoreOperation::create(*addr, *data, states, *response);
516 auto newStore = dynamic_cast<rvsdg::NodeOutput *>(storeOuts[0])->node();
517 // iterate over output states
518 for (size_t i = 0; i < replacedStore->noutputs(); ++i)
519 {
520 // create a buffer to avoid a scenario where the reponse port is blocked because a merge waits
521 // for the store
522 // TODO: It might be better to have memstate merges consume individual tokens instead,, and fire
523 // the output once all inputs have consumed
524 const auto bo = BufferOperation::create(*storeOuts[i], 1, true)[0];
525 smap.insert(originalStore->output(i), bo);
526 replacedStore->output(i)->divert_users(bo);
527 }
528 remove(replacedStore);
529 return dynamic_cast<rvsdg::SimpleNode *>(newStore);
530}
531
532static rvsdg::Output *
534 const rvsdg::LambdaNode * lambda,
535 size_t argumentIndex,
537 const std::vector<rvsdg::Node *> & originalLoadNodes,
538 const std::vector<rvsdg::Node *> & originalStoreNodes,
539 const std::vector<rvsdg::Node *> & originalDecoupledNodes)
540{
541 //
542 // We have the memory operations from the original lambda and need to lookup the corresponding
543 // nodes in the new lambda
544 //
545 std::vector<rvsdg::SimpleNode *> loadNodes;
546 std::vector<std::shared_ptr<const rvsdg::Type>> responseTypes;
547 for (auto loadNode : originalLoadNodes)
548 {
549 auto oldLoadedValue = loadNode->output(0);
550 JLM_ASSERT(smap.contains(*oldLoadedValue));
551 auto & newLoadNode = rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.lookup(*oldLoadedValue));
552 loadNodes.push_back(&newLoadNode);
553 auto loadOp =
554 util::assertedCast<const llvm::LoadNonVolatileOperation>(&newLoadNode.GetOperation());
555 responseTypes.push_back(loadOp->GetLoadedType());
556 }
557 std::vector<rvsdg::SimpleNode *> decoupledNodes;
558 for (auto decoupleRequest : originalDecoupledNodes)
559 {
560 auto oldOutput = decoupleRequest->output(0);
561 JLM_ASSERT(smap.contains(*oldOutput));
562 auto & decoupledRequestNode =
563 rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.lookup(*oldOutput));
564 decoupledNodes.push_back(&decoupledRequestNode);
565 // get load type from response output
566 auto channel = decoupleRequest->input(1)->origin();
567 auto channelConstant = trace_constant(channel);
568 auto reponse = find_decouple_response(lambda, channelConstant);
569 auto vt = reponse->output(0)->Type();
570 responseTypes.push_back(vt);
571 }
572 std::vector<rvsdg::SimpleNode *> storeNodes;
573 for (auto storeNode : originalStoreNodes)
574 {
575 auto oldOutput = storeNode->output(0);
576 JLM_ASSERT(smap.contains(*oldOutput));
577 auto & newStoreNode = rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.lookup(*oldOutput));
578 storeNodes.push_back(&newStoreNode);
579 // use memory state type as response for stores
580 auto vt = std::make_shared<llvm::MemoryStateType>();
581 responseTypes.push_back(vt);
582 }
583
584 auto lambdaRegion = lambda->subregion();
585 auto portWidth =
586 CalculatePortWidth({ originalLoadNodes, originalStoreNodes, originalDecoupledNodes });
587 auto responses = MemoryResponseOperation::create(
588 *lambdaRegion->argument(argumentIndex),
589 responseTypes,
590 portWidth);
591 // The (decoupled) load nodes are replaced so the pointer to the types will become invalid
592 std::vector<std::shared_ptr<const rvsdg::Type>> loadTypes;
593 std::vector<rvsdg::Output *> loadAddresses;
594 for (size_t i = 0; i < loadNodes.size(); ++i)
595 {
596 auto routed = route_response_rhls(loadNodes[i]->region(), responses[i]);
597 // The smap contains the nodes from the original lambda so we need to use the original load node
598 // when replacing the load since the smap must be updated
599 auto replacement = ReplaceLoad(smap, originalLoadNodes[i], routed);
600 auto address =
601 route_request_rhls(lambdaRegion, replacement->output(replacement->noutputs() - 1));
602 loadAddresses.push_back(address);
603 std::shared_ptr<const rvsdg::Type> type;
604 if (auto loadOperation = dynamic_cast<const LoadOperation *>(&replacement->GetOperation()))
605 {
606 type = loadOperation->GetLoadedType();
607 }
608 else if (
609 auto loadOperation =
610 dynamic_cast<const DecoupledLoadOperation *>(&replacement->GetOperation()))
611 {
612 type = loadOperation->GetLoadedType();
613 }
614 else
615 {
616 JLM_UNREACHABLE("Unknown load GetOperation");
617 }
618 JLM_ASSERT(type);
619 loadTypes.push_back(type);
620 }
621 for (size_t i = 0; i < decoupledNodes.size(); ++i)
622 {
623 auto response = responses[loadNodes.size() + i];
624 auto node = decoupledNodes[i];
625
626 // TODO: this beahvior is not completly correct - if a function returns a top-level result from
627 // a decouple it fails and smap translation would be required
628 auto replacement = ReplaceDecouple(lambda, node, response);
629 auto addr = route_request_rhls(lambdaRegion, replacement->output(1));
630 loadAddresses.push_back(addr);
631 loadTypes.push_back(dynamic_cast<const DecoupledLoadOperation *>(&replacement->GetOperation())
632 ->GetLoadedType());
633 }
634 std::vector<rvsdg::Output *> storeOperands;
635 for (size_t i = 0; i < storeNodes.size(); ++i)
636 {
637 auto response = responses[loadNodes.size() + decoupledNodes.size() + i];
638 auto routed = route_response_rhls(storeNodes[i]->region(), response);
639 // The smap contains the nodes from the original lambda so we need to use the original store
640 // node when replacing the store since the smap must be updated
641 auto replacement = ReplaceStore(smap, originalStoreNodes[i], routed);
642 auto addr = route_request_rhls(lambdaRegion, replacement->output(replacement->noutputs() - 2));
643 auto data = route_request_rhls(lambdaRegion, replacement->output(replacement->noutputs() - 1));
644 storeOperands.push_back(addr);
645 storeOperands.push_back(data);
646 }
647
648 return MemoryRequestOperation::create(loadAddresses, loadTypes, storeOperands, lambdaRegion)[0];
649}
650
651static void
653{
654 //
655 // Replacing memory nodes with nodes that have explicit memory ports requires arguments and
656 // results to be added to the lambda. The arguments must be added before the memory nodes are
657 // replaced, else the input of the new memory node will be left dangling, which is not allowed. We
658 // therefore need to first replace the lambda node with a new lambda node that has the new
659 // arguments and results. We can then replace the memory nodes and connect them to the new
660 // arguments.
661 //
662
663 const auto & graph = rvsdgModule.Rvsdg();
664 const auto rootRegion = &graph.GetRootRegion();
665 if (rootRegion->numNodes() != 1)
666 {
667 throw std::logic_error("Root should have only one node now");
668 }
669
670 const auto lambda = dynamic_cast<rvsdg::LambdaNode *>(rootRegion->Nodes().begin().ptr());
671 if (!lambda)
672 {
673 throw std::logic_error("Node needs to be a lambda");
674 }
675
676 //
677 // Converting loads and stores to explicitly use memory ports
678 // This modifies the function signature so we create a new lambda node to replace the old one
679 //
680 const auto & op = dynamic_cast<llvm::LlvmLambdaOperation &>(lambda->GetOperation());
681 auto oldFunctionType = op.type();
682 std::vector<std::shared_ptr<const rvsdg::Type>> newArgumentTypes;
683 for (size_t i = 0; i < oldFunctionType.NumArguments(); ++i)
684 {
685 newArgumentTypes.push_back(oldFunctionType.Arguments()[i]);
686 }
687 std::vector<std::shared_ptr<const rvsdg::Type>> newResultTypes;
688 for (size_t i = 0; i < oldFunctionType.NumResults(); ++i)
689 {
690 newResultTypes.push_back(oldFunctionType.Results()[i]);
691 }
692
693 //
694 // Get the load and store nodes and add an argument and result for each to represent the memory
695 // response and request ports
696 //
697 auto tracedPointerNodesVector = TracePointerArguments(lambda);
698
699 std::unordered_set<rvsdg::Node *> accountedNodes;
700 for (auto & portNode : tracedPointerNodesVector)
701 {
702 if (portNode.isEmpty())
703 continue;
704
705 auto portWidth = CalculatePortWidth(portNode);
706 auto responseTypePtr = get_mem_res_type(rvsdg::BitType::Create(portWidth));
707 auto requestTypePtr = get_mem_req_type(rvsdg::BitType::Create(portWidth), false);
708 auto requestTypePtrWrite = get_mem_req_type(rvsdg::BitType::Create(portWidth), true);
709 newArgumentTypes.push_back(responseTypePtr);
710 if (portNode.storeNodes.empty())
711 {
712 newResultTypes.push_back(requestTypePtr);
713 }
714 else
715 {
716 newResultTypes.push_back(requestTypePtrWrite);
717 }
718 accountedNodes.insert(portNode.loadNodes.begin(), portNode.loadNodes.end());
719 accountedNodes.insert(portNode.storeNodes.begin(), portNode.storeNodes.end());
720 accountedNodes.insert(portNode.decoupleNodes.begin(), portNode.decoupleNodes.end());
721 }
722 std::vector<rvsdg::Node *> unknownLoadNodes;
723 std::vector<rvsdg::Node *> unknownStoreNodes;
724 std::vector<rvsdg::Node *> unknownDecoupledNodes;
726 rootRegion,
727 unknownLoadNodes,
728 unknownStoreNodes,
729 unknownDecoupledNodes,
730 accountedNodes);
731 if (!unknownLoadNodes.empty() || !unknownStoreNodes.empty() || !unknownDecoupledNodes.empty())
732 {
733 auto portWidth =
734 CalculatePortWidth({ unknownLoadNodes, unknownStoreNodes, unknownDecoupledNodes });
735 auto responseTypePtr = get_mem_res_type(rvsdg::BitType::Create(portWidth));
736 auto requestTypePtr = get_mem_req_type(rvsdg::BitType::Create(portWidth), false);
737 auto requestTypePtrWrite = get_mem_req_type(rvsdg::BitType::Create(portWidth), true);
738 // Extra port for loads/stores not associated to a port yet (i.e., unknown base pointer)
739 newArgumentTypes.push_back(responseTypePtr);
740 if (unknownStoreNodes.empty())
741 {
742 newResultTypes.push_back(requestTypePtr);
743 }
744 else
745 {
746 newResultTypes.push_back(requestTypePtrWrite);
747 }
748 }
749
750 //
751 // Create new lambda and copy the region from the old lambda
752 //
753 auto newFunctionType = rvsdg::FunctionType::Create(newArgumentTypes, newResultTypes);
754 auto newLambda = rvsdg::LambdaNode::Create(
755 *lambda->region(),
756 llvm::LlvmLambdaOperation::Create(
757 newFunctionType,
758 op.name(),
759 op.linkage(),
760 op.callingConvention(),
761 op.attributes()));
762
764 for (const auto & ctxvar : lambda->GetContextVars())
765 {
766 smap.insert(ctxvar.inner, newLambda->AddContextVar(*ctxvar.input->origin()).inner);
767 }
768
769 auto args = lambda->GetFunctionArguments();
770 auto newArgs = newLambda->GetFunctionArguments();
771 // The new function has more arguments than the old function.
772 // Substitution of existing arguments is safe, but note
773 // that this is not an isomorphism.
774 JLM_ASSERT(args.size() <= newArgs.size());
775 for (size_t i = 0; i < args.size(); ++i)
776 {
777 smap.insert(args[i], newArgs[i]);
778 }
779 lambda->subregion()->copy(newLambda->subregion(), smap);
780
781 //
782 // All memory nodes need to be replaced with new nodes that have explicit memory ports.
783 // This needs to happen first and the smap needs to be updated with the new nodes,
784 // before we can use the original lambda results and look them up in the updated smap.
785 //
786
787 std::vector<rvsdg::Output *> newResults;
788 // The new arguments are placed directly after the original arguments so we create an index that
789 // points to the first new argument
790 auto newArgumentsIndex = args.size();
791 for (auto & portNode : tracedPointerNodesVector)
792 {
793 if (!portNode.isEmpty())
794 {
795 newResults.push_back(ConnectRequestResponseMemPorts(
796 newLambda,
797 newArgumentsIndex++,
798 smap,
799 portNode.loadNodes,
800 portNode.storeNodes,
801 portNode.decoupleNodes));
802 }
803 }
804 if (!unknownLoadNodes.empty() || !unknownStoreNodes.empty() || !unknownDecoupledNodes.empty())
805 {
806 newResults.push_back(ConnectRequestResponseMemPorts(
807 newLambda,
808 newArgumentsIndex++,
809 smap,
810 unknownLoadNodes,
811 unknownStoreNodes,
812 unknownDecoupledNodes));
813 }
814
815 std::vector<rvsdg::Output *> originalResults;
816 for (auto result : lambda->GetFunctionResults())
817 {
818 originalResults.push_back(&smap.lookup(*result->origin()));
819 }
820 originalResults.insert(originalResults.end(), newResults.begin(), newResults.end());
821 auto newOut = newLambda->finalize(originalResults);
822 auto oldExport = llvm::ComputeCallSummary(*lambda).GetRvsdgExport();
823 rvsdg::GraphExport::Create(*newOut, oldExport ? oldExport->Name() : "");
824
825 JLM_ASSERT(lambda->output()->nusers() == 1);
826 lambda->region()->RemoveResults({ (*lambda->output()->Users().begin()).index() });
827 remove(lambda);
828
829 // Remove imports for decouple_ function pointers
832 dne.Run(*newLambda->subregion(), statisticsCollector);
833
834 //
835 // TODO
836 // RemoveUnusedStates also creates a new lambda, which we have already done above.
837 // It might be better to apply this functionality above such that we only create a new lambda
838 // once.
839 //
840 UnusedStateRemoval::CreateAndRun(rvsdgModule, statisticsCollector);
841
842 // Need to get the lambda from the root since remote_unused_state replaces the lambda
843 JLM_ASSERT(rootRegion->numNodes() == 1);
844 newLambda = util::assertedCast<rvsdg::LambdaNode>(rootRegion->Nodes().begin().ptr());
845 auto decouple_funcs = find_function_arguments(newLambda, "decoupled");
846 // make sure context vars are actually dead
847 for (auto cv : decouple_funcs)
848 {
849 JLM_ASSERT(cv.inner->nusers() == 0);
850 }
851 // remove dead cvargs
852 newLambda->PruneLambdaInputs();
853}
854
855MemoryConverter::~MemoryConverter() noexcept = default;
856
860
861void
862MemoryConverter::Run(rvsdg::RvsdgModule & rvsdgModule, util::StatisticsCollector &)
863{
864 ConvertMemory(rvsdgModule);
865}
866
867}
static const auto vt
Definition PullTests.cpp:16
static jlm::util::StatisticsCollector statisticsCollector
Definition PullTests.cpp:17
static std::vector< jlm::rvsdg::Output * > create(jlm::rvsdg::Output &addr, jlm::rvsdg::Output &load_result, size_t capacity)
Definition hls.hpp:1218
std::shared_ptr< const rvsdg::Type > GetLoadedType() const noexcept
Definition hls.hpp:1235
void Run(rvsdg::RvsdgModule &rvsdgModule, util::StatisticsCollector &statisticsCollector) override
Perform RVSDG transformation.
Definition rhls-dne.cpp:512
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 rvsdg::Output * Create(const std::vector< rvsdg::Output * > &operands)
static jlm::rvsdg::Output * Create(rvsdg::Region &region, std::shared_ptr< const jlm::rvsdg::Type > type)
Conditional operator / pattern matching.
Definition gamma.hpp:99
Region & GetRootRegion() const noexcept
Definition graph.hpp:99
std::vector< rvsdg::Output * > GetFunctionArguments() const
Definition lambda.cpp:58
rvsdg::Region * subregion() const noexcept
Definition lambda.hpp:138
std::vector< ContextVar > GetContextVars() const noexcept
Gets all bound context variables.
Definition lambda.cpp:120
const FunctionType & type() const noexcept
Definition lambda.hpp:36
NodeOutput * output(size_t index) const noexcept
Definition node.hpp:650
rvsdg::Region * region() const noexcept
Definition node.hpp:761
UsersRange Users()
Definition node.hpp:354
const std::shared_ptr< const rvsdg::Type > & Type() const noexcept
Definition node.hpp:366
Represents the result of a region.
Definition region.hpp:120
Represent acyclic RVSDG subgraphs.
Definition region.hpp:213
rvsdg::StructuralNode * node() const noexcept
Definition region.hpp:301
Graph & Rvsdg() noexcept
void insert(const Output *original, Output *substitute)
Output & lookup(const Output &original) const
bool contains(const Output &original) const noexcept
Represents an RVSDG transformation.
#define JLM_ASSERT(x)
Definition common.hpp:16
#define JLM_UNREACHABLE(msg)
Definition common.hpp:43
rvsdg::LambdaNode * find_containing_lambda(rvsdg::Region *region)
Definition mem-conv.cpp:415
std::shared_ptr< const BundleType > get_mem_res_type(std::shared_ptr< const jlm::rvsdg::Type > dataType)
Definition hls.cpp:448
rvsdg::Output * route_response_rhls(rvsdg::Region *target, rvsdg::Output *response)
void gather_mem_nodes(rvsdg::Region *region, std::vector< rvsdg::Node * > &loadNodes, std::vector< rvsdg::Node * > &storeNodes, std::vector< rvsdg::Node * > &decoupleNodes, std::unordered_set< rvsdg::Node * > exclude)
Definition mem-conv.cpp:228
std::shared_ptr< const BundleType > get_mem_req_type(std::shared_ptr< const rvsdg::Type > elementType, bool write)
Definition hls.cpp:433
void trace_function_calls(rvsdg::Output *output, std::vector< rvsdg::SimpleNode * > &calls, std::unordered_set< rvsdg::Output * > &visited)
bool is_function_argument(const rvsdg::LambdaNode::ContextVar &cv)
rvsdg::Output * route_request_rhls(rvsdg::Region *target, rvsdg::Output *request)
rvsdg::SimpleNode * ReplaceDecouple(const rvsdg::LambdaNode *lambda, rvsdg::SimpleNode *decouple_request, rvsdg::Output *resp)
Definition mem-conv.cpp:144
static void TracePointer(rvsdg::Output *output, std::unordered_set< rvsdg::Output * > &visited, TracedPointerNodes &tracedPointerNodes)
Definition mem-conv.cpp:274
int JlmSize(const jlm::rvsdg::Type *type)
Definition hls.cpp:457
rvsdg::SimpleNode * find_decouple_response(const rvsdg::LambdaNode *lambda, const llvm::IntegerConstantOperation *request_constant)
Definition mem-conv.cpp:30
rvsdg::Output * route_to_region_rhls(rvsdg::Region *target, rvsdg::Output *out)
void OptimizeReqMemState(rvsdg::Output *req_mem_state)
Definition mem-conv.cpp:122
static void ConvertMemory(rvsdg::RvsdgModule &rvsdgModule)
Definition mem-conv.cpp:652
const llvm::IntegerConstantOperation * trace_constant(const rvsdg::Output *dst)
std::vector< TracedPointerNodes > TracePointerArguments(const rvsdg::LambdaNode *lambda)
Definition mem-conv.cpp:388
static rvsdg::SimpleNode * ReplaceStore(rvsdg::SubstitutionMap &smap, const rvsdg::Node *originalStore, rvsdg::Output *response)
Definition mem-conv.cpp:496
static rvsdg::Output * ConnectRequestResponseMemPorts(const rvsdg::LambdaNode *lambda, size_t argumentIndex, rvsdg::SubstitutionMap &smap, const std::vector< rvsdg::Node * > &originalLoadNodes, const std::vector< rvsdg::Node * > &originalStoreNodes, const std::vector< rvsdg::Node * > &originalDecoupledNodes)
Definition mem-conv.cpp:533
static std::pair< rvsdg::Input *, std::vector< rvsdg::Input * > > TraceEdgeToMerge(rvsdg::Input *state_edge)
Definition mem-conv.cpp:53
rvsdg::Input * get_mem_state_user(rvsdg::Output *state_edge)
static size_t CalculatePortWidth(const TracedPointerNodes &tracedPointerNodes)
Definition mem-conv.cpp:425
static rvsdg::SimpleNode * ReplaceLoad(rvsdg::SubstitutionMap &smap, const rvsdg::Node *originalLoad, rvsdg::Output *response)
Definition mem-conv.cpp:454
std::vector< rvsdg::LambdaNode::ContextVar > find_function_arguments(const rvsdg::LambdaNode *lambda, std::string name_contains)
bool is_dec_req(rvsdg::SimpleNode *node)
void OptimizeResMemState(rvsdg::Output *res_mem_state)
Definition mem-conv.cpp:102
static void remove(Node *node)
Definition region.hpp:1035
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
detail::TopDownTraverserGeneric< false > TopDownTraverser
Traverser for visiting every node in a region in a top down order.
std::vector< rvsdg::Node * > loadNodes
Definition mem-conv.hpp:24
std::vector< rvsdg::Node * > decoupleNodes
Definition mem-conv.hpp:26
std::vector< rvsdg::Node * > storeNodes
Definition mem-conv.hpp:25
A variable routed into all gamma regions.
Definition gamma.hpp:131