276 std::unordered_set<rvsdg::Output *> & visited,
284 if (visited.count(output))
289 visited.insert(output);
296 rvsdg::MatchTypeOrFail(
298 [&](rvsdg::SimpleNode & simplenode)
300 rvsdg::MatchTypeWithDefault(
301 simplenode.GetOperation(),
302 [&](const llvm::StoreNonVolatileOperation &)
304 tracedPointerNodes.storeNodes.push_back(&simplenode);
306 [&](const llvm::LoadNonVolatileOperation &)
308 tracedPointerNodes.loadNodes.push_back(&simplenode);
310 [&](const llvm::CallOperation &)
312 JLM_ASSERT(is_dec_req(&simplenode));
313 tracedPointerNodes.decoupleNodes.push_back(&simplenode);
317 for (size_t i = 0; i < simplenode.noutputs(); ++i)
319 TracePointer(simplenode.output(i), visited, tracedPointerNodes);
325 TracePointer(loop.mapInput(user).inner, visited, tracedPointerNodes);
329 TracePointer(theta.MapInputLoopVar(user).pre, visited, tracedPointerNodes);
334 gamma.MapInput(user),
335 [&](const rvsdg::GammaNode::MatchVar &)
340 for (auto arg : evar.branchArgument)
342 TracePointer(arg, visited, tracedPointerNodes);
347 [&](rvsdg::Region * region)
349 rvsdg::MatchTypeOrFail(
354 loop.mapResult(user),
355 [&](const LoopNode::BackEdgeVar & backedge)
357 TracePointer(backedge.pre, visited, tracedPointerNodes);
359 [&](
const LoopNode::ExitVar & exit)
361 TracePointer(exit.output, visited, tracedPointerNodes);
364 [&](rvsdg::ThetaNode & theta)
367 theta.mapResult(user),
368 [&](
const rvsdg::ThetaNode::LoopVar & loopvar)
370 TracePointer(loopvar.output, visited, tracedPointerNodes);
372 [&](
const rvsdg::ThetaNode::PredicateVar &)
376 [&](rvsdg::GammaNode & gamma)
379 gamma.MapBranchResultExitVar(user).output,
428 for (
auto node : tracedPointerNodes.
loadNodes)
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;
434 for (
auto node : tracedPointerNodes.
storeNodes)
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;
440 for (
auto decoupleRequest : tracedPointerNodes.
decoupleNodes)
443 auto channel = decoupleRequest->input(1)->origin();
446 auto sz =
JlmSize(reponse->output(0)->Type().get());
447 max_width = sz > max_width ? sz : max_width;
505 &rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.
lookup(*originalStore->
output(0)));
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)
513 states.push_back(replacedStore->input(i)->origin());
515 auto storeOuts = StoreOperation::create(*addr, *data, states, *response);
518 for (
size_t i = 0; i < replacedStore->noutputs(); ++i)
524 const auto bo = BufferOperation::create(*storeOuts[i], 1,
true)[0];
526 replacedStore->output(i)->divert_users(bo);
528 remove(replacedStore);
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)
545 std::vector<rvsdg::SimpleNode *> loadNodes;
546 std::vector<std::shared_ptr<const rvsdg::Type>> responseTypes;
547 for (
auto loadNode : originalLoadNodes)
549 auto oldLoadedValue = loadNode->output(0);
551 auto & newLoadNode = rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.
lookup(*oldLoadedValue));
552 loadNodes.push_back(&newLoadNode);
554 util::assertedCast<const llvm::LoadNonVolatileOperation>(&newLoadNode.GetOperation());
555 responseTypes.push_back(loadOp->GetLoadedType());
557 std::vector<rvsdg::SimpleNode *> decoupledNodes;
558 for (
auto decoupleRequest : originalDecoupledNodes)
560 auto oldOutput = decoupleRequest->output(0);
562 auto & decoupledRequestNode =
563 rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.
lookup(*oldOutput));
564 decoupledNodes.push_back(&decoupledRequestNode);
566 auto channel = decoupleRequest->input(1)->origin();
569 auto vt = reponse->output(0)->Type();
570 responseTypes.push_back(
vt);
572 std::vector<rvsdg::SimpleNode *> storeNodes;
573 for (
auto storeNode : originalStoreNodes)
575 auto oldOutput = storeNode->output(0);
577 auto & newStoreNode = rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(smap.
lookup(*oldOutput));
578 storeNodes.push_back(&newStoreNode);
580 auto vt = std::make_shared<llvm::MemoryStateType>();
581 responseTypes.push_back(
vt);
586 CalculatePortWidth({ originalLoadNodes, originalStoreNodes, originalDecoupledNodes });
587 auto responses = MemoryResponseOperation::create(
588 *lambdaRegion->argument(argumentIndex),
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)
599 auto replacement =
ReplaceLoad(smap, originalLoadNodes[i], routed);
602 loadAddresses.push_back(address);
603 std::shared_ptr<const rvsdg::Type> type;
604 if (
auto loadOperation =
dynamic_cast<const LoadOperation *
>(&replacement->GetOperation()))
606 type = loadOperation->GetLoadedType();
612 type = loadOperation->GetLoadedType();
619 loadTypes.push_back(type);
621 for (
size_t i = 0; i < decoupledNodes.size(); ++i)
623 auto response = responses[loadNodes.size() + i];
624 auto node = decoupledNodes[i];
630 loadAddresses.push_back(addr);
634 std::vector<rvsdg::Output *> storeOperands;
635 for (
size_t i = 0; i < storeNodes.size(); ++i)
637 auto response = responses[loadNodes.size() + decoupledNodes.size() + i];
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);
648 return MemoryRequestOperation::create(loadAddresses, loadTypes, storeOperands, lambdaRegion)[0];
663 const auto & graph = rvsdgModule.
Rvsdg();
665 if (rootRegion->numNodes() != 1)
667 throw std::logic_error(
"Root should have only one node now");
670 const auto lambda =
dynamic_cast<rvsdg::LambdaNode *
>(rootRegion->Nodes().begin().ptr());
673 throw std::logic_error(
"Node needs to be a lambda");
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)
685 newArgumentTypes.push_back(oldFunctionType.Arguments()[i]);
687 std::vector<std::shared_ptr<const rvsdg::Type>> newResultTypes;
688 for (
size_t i = 0; i < oldFunctionType.NumResults(); ++i)
690 newResultTypes.push_back(oldFunctionType.Results()[i]);
699 std::unordered_set<rvsdg::Node *> accountedNodes;
700 for (
auto & portNode : tracedPointerNodesVector)
702 if (portNode.isEmpty())
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())
712 newResultTypes.push_back(requestTypePtr);
716 newResultTypes.push_back(requestTypePtrWrite);
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());
722 std::vector<rvsdg::Node *> unknownLoadNodes;
723 std::vector<rvsdg::Node *> unknownStoreNodes;
724 std::vector<rvsdg::Node *> unknownDecoupledNodes;
729 unknownDecoupledNodes,
731 if (!unknownLoadNodes.empty() || !unknownStoreNodes.empty() || !unknownDecoupledNodes.empty())
736 auto requestTypePtr =
get_mem_req_type(rvsdg::BitType::Create(portWidth),
false);
737 auto requestTypePtrWrite =
get_mem_req_type(rvsdg::BitType::Create(portWidth),
true);
739 newArgumentTypes.push_back(responseTypePtr);
740 if (unknownStoreNodes.empty())
742 newResultTypes.push_back(requestTypePtr);
746 newResultTypes.push_back(requestTypePtrWrite);
753 auto newFunctionType = rvsdg::FunctionType::Create(newArgumentTypes, newResultTypes);
754 auto newLambda = rvsdg::LambdaNode::Create(
756 llvm::LlvmLambdaOperation::Create(
760 op.callingConvention(),
764 for (
const auto & ctxvar : lambda->GetContextVars())
766 smap.
insert(ctxvar.inner, newLambda->AddContextVar(*ctxvar.input->origin()).inner);
769 auto args = lambda->GetFunctionArguments();
770 auto newArgs = newLambda->GetFunctionArguments();
775 for (
size_t i = 0; i < args.size(); ++i)
777 smap.
insert(args[i], newArgs[i]);
779 lambda->subregion()->copy(newLambda->subregion(), smap);
787 std::vector<rvsdg::Output *> newResults;
790 auto newArgumentsIndex = args.size();
791 for (
auto & portNode : tracedPointerNodesVector)
793 if (!portNode.isEmpty())
801 portNode.decoupleNodes));
804 if (!unknownLoadNodes.empty() || !unknownStoreNodes.empty() || !unknownDecoupledNodes.empty())
812 unknownDecoupledNodes));
815 std::vector<rvsdg::Output *> originalResults;
816 for (
auto result : lambda->GetFunctionResults())
818 originalResults.push_back(&smap.
lookup(*result->origin()));
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() :
"");
826 lambda->region()->RemoveResults({ (*lambda->output()->Users().begin()).index() });
844 newLambda = util::assertedCast<rvsdg::LambdaNode>(rootRegion->Nodes().begin().ptr());
847 for (
auto cv : decouple_funcs)
852 newLambda->PruneLambdaInputs();
A variable routed into all gamma regions.