98 static std::unique_ptr<Statistics>
101 return std::make_unique<Statistics>(sourceFile);
137 static std::unique_ptr<Context>
140 return std::make_unique<Context>();
150 : Transformation(
"IOBarrierElimination")
158 auto & rvsdg =
module.Rvsdg();
163 statistics->startNormalizationStatistics();
165 statistics->stopNormalizationStatistics();
167 statistics->startMarkStatistics();
169 statistics->stopMarkStatistics();
171 statistics->startPropagateStatistics();
173 statistics->stopPropagateStatistics();
175 statistics->startEncodeStatistics();
177 statistics->stopEncodeStatistics();
179 statistics->startSweepStatistics();
181 statistics->stopSweepStatistics();
189static std::vector<rvsdg::SimpleNode *>
192 std::vector<rvsdg::SimpleNode *> hoistBarrierNodes;
193 for (
auto & user : output.
Users())
195 if (
auto [node, hoistBarrierOp] =
199 hoistBarrierNodes.push_back(node);
203 return hoistBarrierNodes;
206static std::optional<rvsdg::SimpleNode *>
208 const std::vector<rvsdg::SimpleNode *> & hoistBarrierNodes,
209 const std::variant<rvsdg::Node *, rvsdg::Region *> ioStateOwner)
211 for (
auto node : hoistBarrierNodes)
225 JLM_ASSERT(is<MemoryHoistBarrierOperation>(&memoryHoistBarrierNode));
228 *memoryHoistBarrierNode.
output(0),
231 return &MemoryHoistBarrierOperation::getAddressInput(memoryHoistBarrierNode) != &user;
235static std::optional<rvsdg::GammaNode::EntryVar>
250 for (
auto & node : region.
Nodes())
256 for (
auto & subregion : structuralNode.
Subregions())
262 for (
auto & argument : subregion.Arguments())
264 if (is<PointerType>(argument->Type()))
274 for (
auto & output : structuralNode.
Outputs())
276 if (is<PointerType>(output.Type()))
279 if (
auto ioBarrierNode =
288 simpleNode.GetOperation(),
291 auto & loadedValue = LoadOperation::LoadedValueOutput(simpleNode);
292 if (is<PointerType>(loadedValue.Type()))
294 auto ioBarrierNodes = collectMemoryHoistBarrierNodes(loadedValue);
295 if (auto ioBarrierNode =
296 selectMemoryHoistBarrierNode(ioBarrierNodes, simpleNode.region()))
297 divertUsersToMemoryHoistBarrierNode(loadedValue, **ioBarrierNode);
303 throw std::logic_error(
"Unexpected node type");
315 size_t size = std::numeric_limits<std::size_t>::max();
318 if (
const size_t numUsers = argument->nusers(); numUsers == 0)
324 if (
auto argumentSize = context_->getDereferenceableSize(*argument); argumentSize > 0)
328 size = std::min(size, argumentSize);
334 for (
auto & user : argument->Users())
336 auto [simpleNode, hoistBarrierOp] =
341 auto & ioStateOperand =
342 *MemoryHoistBarrierOperation::getIOStateInput(*simpleNode).origin();
349 hoistBarrierNode = simpleNode;
354 if (!hoistBarrierNode)
360 argumentSize = context_->getDereferenceableSize(
361 MemoryHoistBarrierOperation::getAddressOutput(*hoistBarrierNode));
362 if (argumentSize == 0)
371 size = std::min(size, argumentSize);
381 for (
auto & node : region.
Nodes())
392 for (
const auto argument : lambdaNode.GetFunctionArguments())
396 context_->markDereferenceable(*argument, 0);
401 for (
const auto [_, inner] : lambdaNode.GetContextVars())
406 context_->markDereferenceable(*inner, 0);
410 markOutputs(*lambdaNode.subregion());
418 markOutputs(*thetaNode.subregion());
423 for (
auto & subregion : gammaNode.Subregions())
425 markOutputs(subregion);
430 for (
auto & entryVar : gammaNode.GetEntryVars())
432 if (
const auto size = getDereferenceableSize(entryVar); size > 0)
438 auto & mhbNode = MemoryHoistBarrierOperation::createNode(
439 *entryVar.input->origin(),
440 *ioStateEntryVar->input->origin(),
442 auto & mhbAddressOutput = MemoryHoistBarrierOperation::getAddressOutput(mhbNode);
443 entryVar.input->divert_to(&mhbAddressOutput);
444 context_->markDereferenceable(mhbAddressOutput, size);
452 simpleNode.GetOperation(),
455 const auto & addressOperand = *LoadOperation::AddressInput(simpleNode).origin();
456 const auto sizeInBytes = GetTypeStoreSize(*loadOperation.GetLoadedType());
457 context_->markDereferenceable(addressOperand, sizeInBytes);
461 const auto & addressOperand = *StoreOperation::AddressInput(simpleNode).origin();
462 const auto sizeInBytes = GetTypeStoreSize(storeOperation.GetStoredType());
463 context_->markDereferenceable(addressOperand, sizeInBytes);
468 throw std::logic_error(
"Unhandled node type");
488 propagate(*lambdaNode.subregion());
492 for (
auto & [input,
arguments] : gammaNode.GetEntryVars())
494 if (!is<PointerType>(input->Type()))
497 if (
const auto size = context_->getDereferenceableSize(*input->origin()); size > 0)
501 context_->markDereferenceable(*argument, size);
506 for (
auto & subregion : gammaNode.Subregions())
507 propagate(subregion);
509 for (
auto & [results, output] : gammaNode.GetExitVars())
511 if (!is<PointerType>(output->Type()))
514 size_t sizeInBytes = std::numeric_limits<std::size_t>::max();
515 for (
const auto & result : results)
518 std::min(sizeInBytes, context_->getDereferenceableSize(*result->origin()));
519 if (sizeInBytes == 0)
525 context_->markDereferenceable(*output, sizeInBytes);
532 std::unordered_map<rvsdg::Output *, size_t> loopVarPreSizes;
535 for (
const auto & loopVar : thetaNode.GetLoopVars())
537 if (!is<PointerType>(loopVar.input->Type()))
540 if (
const auto inputSize = context_->getDereferenceableSize(*loopVar.input->origin());
543 loopVarPreSizes[loopVar.pre] = inputSize;
544 context_->markDereferenceable(*loopVar.pre, inputSize);
553 propagate(*thetaNode.subregion());
555 for (
const auto & loopVar : thetaNode.GetLoopVars())
557 if (!is<PointerType>(loopVar.input->Type()))
560 const auto preSize = loopVarPreSizes[loopVar.pre];
561 const auto postSize = context_->getDereferenceableSize(*loopVar.post->origin());
562 if (preSize != postSize)
564 loopVarPreSizes[loopVar.pre] = postSize;
565 context_->markDereferenceable(*loopVar.pre, std::min(preSize, postSize));
572 for (
const auto & loopVar : thetaNode.GetLoopVars())
574 if (!is<PointerType>(loopVar.output->Type()))
577 if (
const auto postSize = context_->getDereferenceableSize(*loopVar.post->origin());
580 context_->markDereferenceable(*loopVar.output, postSize);
591 simpleNode.GetOperation(),
594 const auto & barredInput =
595 MemoryHoistBarrierOperation::getAddressInput(simpleNode);
596 if (!is<PointerType>(barredInput.Type()))
599 if (const auto size = context_->getDereferenceableSize(*barredInput.origin());
601 context_->markDereferenceable(*simpleNode.output(0), size);
606 throw std::logic_error(
607 "Unhandled node type encountered during dereferenceable propagation.");
630 encode(*lambdaNode.subregion());
634 for (
auto & subregion : gammaNode.Subregions())
639 encode(*thetaNode.subregion());
648 simpleNode.GetOperation(),
651 auto & addressOperand =
652 *MemoryHoistBarrierOperation::getAddressInput(simpleNode).origin();
653 const auto mhbSize = memoryHoistBarrier.getDereferenceableSize();
655 if (const auto size = context_->getDereferenceableSize(addressOperand);
658 auto & ioStateOperand =
659 *MemoryHoistBarrierOperation::getIOStateInput(simpleNode).origin();
660 auto & mhbNode = MemoryHoistBarrierOperation::createNode(
664 std::max(size, mhbSize));
665 MemoryHoistBarrierOperation::getAddressOutput(simpleNode)
666 .divert_users(&MemoryHoistBarrierOperation::getAddressOutput(mhbNode));
672 throw std::logic_error(
"Unhandled node type encountered during encoding.");
679 encode(graph.GetRootRegion());
685 for (
auto & node : region.
Nodes())
695 sweepRegion(*lambdaNode.subregion());
703 for (
auto & subregion : gammaNode.Subregions())
705 sweepRegion(subregion);
710 sweepRegion(*thetaNode.subregion());
714 if (
const auto loadOperation =
717 auto & loadAddress = LoadOperation::AddressInput(simpleNode);
718 auto [hoistBarrierNode, hoistBarrierOp] =
720 *loadAddress.origin());
724 auto & barredAddressInput =
725 MemoryHoistBarrierOperation::getAddressInput(*hoistBarrierNode);
726 const auto barredAddressSize =
727 context_->getDereferenceableSize(*barredAddressInput.origin());
728 if (barredAddressSize == 0)
731 if (
const auto storeSize =
GetTypeStoreSize(*loadOperation->GetLoadedType());
732 barredAddressSize < storeSize)
735 loadAddress.divert_to(barredAddressInput.origin());
740 throw std::logic_error(
"Unsupported node type");
util::HashSet< rvsdg::Output * > arguments
std::unordered_map< const rvsdg::Output *, size_t > dereferenceableInputs_
size_t getDereferenceableSize(const rvsdg::Output &output) const
void markDereferenceable(const rvsdg::Output &output, const size_t sizeInBytes)
static std::unique_ptr< Context > create()
void startEncodeStatistics() noexcept
~Statistics() override=default
Statistics(const util::FilePath &sourceFile)
void startNormalizationStatistics() noexcept
void stopSweepStatistics() noexcept
void stopMarkStatistics() noexcept
const char * NormalizationTimerLabel_
void startPropagateStatistics() noexcept
const char * SweepTimerLabel_
static std::unique_ptr< Statistics > create(const util::FilePath &sourceFile)
const char * PropagateTimerLabel_
void stopPropagateStatistics() noexcept
void startSweepStatistics() noexcept
void stopEncodeStatistics() noexcept
void startMarkStatistics() noexcept
const char * MarkTimerLabel_
void stopNormalizationStatistics() noexcept
const char * EncodeTimerLabel_
void propagateSize(rvsdg::Graph &graph)
void Run(rvsdg::RvsdgModule &module, util::StatisticsCollector &statisticsCollector) override
Perform RVSDG transformation.
static void normalizeMemoryHoistBarriers(rvsdg::Region ®ion)
void encodeSize(rvsdg::Graph &graph)
void markOutputs(rvsdg::Region ®ion)
~IOBarrierElimination() override
void sweepRegion(rvsdg::Region ®ion)
std::unique_ptr< Context > context_
static rvsdg::Input & getIOStateInput(const rvsdg::Node &node) noexcept
Conditional operator / pattern matching.
std::vector< EntryVar > GetEntryVars() const
Gets all entry variables for this gamma.
Region & GetRootRegion() const noexcept
OutputIteratorRange Outputs() noexcept
size_t divertUsersWhere(Output &newOrigin, const F &match)
std::variant< Node *, Region * > GetOwner() const noexcept
A phi node represents the fixpoint of mutually recursive definitions.
rvsdg::Region * subregion() const noexcept
Represent acyclic RVSDG subgraphs.
void prune(bool recursive)
NodeRange Nodes() noexcept
const std::optional< util::FilePath > & SourceFilePath() const noexcept
NodeOutput * output(size_t index) const noexcept
SubregionIteratorRange Subregions()
void CollectDemandedStatistics(std::unique_ptr< Statistics > statistics)
util::Timer & GetTimer(const std::string &name)
util::Timer & AddTimer(std::string name)
Global memory state passed between functions.
static util::StatisticsCollector statisticsCollector
static std::optional< rvsdg::GammaNode::EntryVar > getIOStateEntryVar(const rvsdg::GammaNode &gammaNode)
static void divertUsersToMemoryHoistBarrierNode(rvsdg::Output &output, rvsdg::SimpleNode &memoryHoistBarrierNode)
size_t GetTypeStoreSize(const rvsdg::Type &type)
static std::optional< rvsdg::SimpleNode * > selectMemoryHoistBarrierNode(const std::vector< rvsdg::SimpleNode * > &hoistBarrierNodes, const std::variant< rvsdg::Node *, rvsdg::Region * > ioStateOwner)
static std::vector< rvsdg::SimpleNode * > collectMemoryHoistBarrierNodes(rvsdg::Output &output)
void MatchTypeWithDefault(T &obj, const Fns &... fns)
Pattern match over subclass type of given object with default handler.
void MatchType(T &obj, const Fns &... fns)
Pattern match over subclass type of given object.
Region * TryGetOwnerRegion(const rvsdg::Input &input) noexcept
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.
detail::TopDownTraverserGeneric< false > TopDownTraverser
Traverser for visiting every node in a region in a top down order.
A variable routed into all gamma regions.
rvsdg::Input * input
Variable at entry point (input to gamma node).
std::vector< rvsdg::Output * > branchArgument
Variable inside each of the branch regions (argument per subregion).