28TEST(NodeHoistingTests, simpleGamma)
34 const auto controlType = ControlType::Create(2);
35 const auto valueType = TestType::createValueType();
36 const auto functionType = FunctionType::Create(
44 auto & rvsdg = rvsdgModule.Rvsdg();
46 auto lambdaNode = LambdaNode::Create(
47 rvsdg.GetRootRegion(),
49 auto controlArgument = lambdaNode->GetFunctionArguments()[0];
50 auto valueArgument = lambdaNode->GetFunctionArguments()[1];
52 auto gammaNode = GammaNode::create(controlArgument, 2);
53 auto entryVar = gammaNode->AddEntryVar(valueArgument);
56 auto constantNode = TestOperation::createNode(gammaNode->subregion(0), {}, { valueType });
57 auto binaryNode = TestOperation::createNode(
58 gammaNode->subregion(0),
59 { entryVar.branchArgument[0], constantNode->output(0) },
63 auto unaryNode = TestOperation::createNode(
64 gammaNode->subregion(1),
65 { entryVar.branchArgument[1] },
68 auto exitVar = gammaNode->AddExitVar({ binaryNode->output(0), unaryNode->output(0) });
70 auto lambdaOutput = lambdaNode->finalize({ exitVar.output });
72 GraphExport::Create(*lambdaOutput,
"x");
85 EXPECT_EQ(lambdaNode->subregion()->numNodes(), 4u);
88 EXPECT_EQ(gammaNode->subregion(0)->numNodes(), 0u);
89 EXPECT_EQ(gammaNode->subregion(1)->numNodes(), 0u);
92TEST(NodeHoistingTests, nestedGamma)
98 const auto controlType = ControlType::Create(2);
99 const auto valueType = TestType::createValueType();
100 const auto functionType = FunctionType::Create(
108 auto & rvsdg = rvsdgModule.Rvsdg();
110 auto lambdaNode = LambdaNode::Create(
111 rvsdg.GetRootRegion(),
113 auto controlArgument = lambdaNode->GetFunctionArguments()[0];
114 auto valueArgument = lambdaNode->GetFunctionArguments()[1];
116 auto gammaNode1 = GammaNode::create(controlArgument, 2);
117 auto controlEntryVar = gammaNode1->AddEntryVar(controlArgument);
118 auto valueEntryVar1 = gammaNode1->AddEntryVar(valueArgument);
121 auto constantNode1 = TestOperation::createNode(gammaNode1->subregion(0), {}, { valueType });
123 auto gammaNode2 = GammaNode::create(controlEntryVar.branchArgument[0], 2);
124 auto valueEntryVar2 = gammaNode2->AddEntryVar(valueEntryVar1.branchArgument[0]);
125 auto valueEntryVar3 = gammaNode2->AddEntryVar(constantNode1->output(0));
128 auto binaryNode = TestOperation::createNode(
129 gammaNode1->subregion(0),
130 { valueEntryVar2.branchArgument[0], valueEntryVar3.branchArgument[0] },
134 auto unaryNode = TestOperation::createNode(
135 gammaNode1->subregion(1),
136 { valueEntryVar2.branchArgument[1] },
139 auto exitVar1 = gammaNode2->AddExitVar({ binaryNode->output(0), unaryNode->output(0) });
142 auto constantNode2 = TestOperation::createNode(gammaNode1->subregion(1), {}, { valueType });
144 auto exitVar2 = gammaNode1->AddExitVar({ exitVar1.output, constantNode2->output(0) });
146 auto lambdaOutput = lambdaNode->finalize({ exitVar2.output });
148 GraphExport::Create(*lambdaOutput,
"x");
161 EXPECT_EQ(lambdaNode->subregion()->numNodes(), 5u);
164 EXPECT_EQ(gammaNode1->subregion(0)->numNodes(), 1u);
165 EXPECT_EQ(gammaNode1->subregion(1)->numNodes(), 0u);
168 EXPECT_EQ(gammaNode2->subregion(0)->numNodes(), 0u);
169 EXPECT_EQ(gammaNode2->subregion(1)->numNodes(), 0u);
172TEST(NodeHoistingTests, simpleTheta)
178 auto controlType = ControlType::Create(2);
179 const auto valueType = TestType::createValueType();
180 const auto functionType = FunctionType::Create(
188 auto & rvsdg = rvsdgModule.Rvsdg();
190 auto lambdaNode = LambdaNode::Create(
191 rvsdg.GetRootRegion(),
193 auto controlArgument = lambdaNode->GetFunctionArguments()[0];
194 auto valueArgument = lambdaNode->GetFunctionArguments()[1];
196 auto thetaNode = ThetaNode::create(lambdaNode->subregion());
198 auto lv1 = thetaNode->AddLoopVar(controlArgument);
199 auto lv2 = thetaNode->AddLoopVar(valueArgument);
200 auto lv3 = thetaNode->AddLoopVar(valueArgument);
201 auto lv4 = thetaNode->AddLoopVar(valueArgument);
203 auto node1 = TestOperation::createNode(thetaNode->subregion(), {}, { valueType });
204 auto node2 = TestOperation::createNode(
205 thetaNode->subregion(),
206 { node1->output(0), lv3.pre },
208 auto node3 = TestOperation::createNode(
209 thetaNode->subregion(),
210 { lv2.pre, node2->output(0) },
213 TestOperation::createNode(thetaNode->subregion(), { lv3.pre, lv4.pre }, { valueType });
215 lv2.post->divert_to(node3->output(0));
216 lv4.post->divert_to(node4->output(0));
218 thetaNode->set_predicate(lv1.pre);
220 lambdaNode->finalize({ thetaNode->output(1) });
233 EXPECT_EQ(lambdaNode->subregion()->numNodes(), 3u);
234 EXPECT_EQ(thetaNode->subregion()->numNodes(), 2u);
236 EXPECT_EQ(lv2.post->origin(), node3->output(0));
237 EXPECT_EQ(lv4.post->origin(), node4->output(0));
240TEST(NodeHoistingTests, invariantMemoryOperation)
248 const auto controlType = ControlType::Create(2);
249 const auto valueType = TestType::createValueType();
250 const auto functionType = FunctionType::Create(
251 { controlType, pointerType, valueType, memoryStateType },
252 { memoryStateType });
255 auto & rvsdg = rvsdgModule.Rvsdg();
257 auto lambdaNode = LambdaNode::Create(
258 rvsdg.GetRootRegion(),
260 auto controlArgument = lambdaNode->GetFunctionArguments()[0];
261 auto pointerArgument = lambdaNode->GetFunctionArguments()[1];
262 auto valueArgument = lambdaNode->GetFunctionArguments()[2];
263 auto memoryStateArgument = lambdaNode->GetFunctionArguments()[3];
265 auto thetaNode = ThetaNode::create(lambdaNode->subregion());
267 auto lvc = thetaNode->AddLoopVar(controlArgument);
268 auto lva = thetaNode->AddLoopVar(pointerArgument);
269 auto lvv = thetaNode->AddLoopVar(valueArgument);
270 auto lvs = thetaNode->AddLoopVar(memoryStateArgument);
274 lvs.post->divert_to(storeNode.output(0));
275 thetaNode->set_predicate(lvc.pre);
277 lambdaNode->finalize({ lvs.output });
290 EXPECT_EQ(lambdaNode->subregion()->numNodes(), 2u);
291 EXPECT_EQ(thetaNode->subregion()->numNodes(), 0u);
295 EXPECT_EQ(thetaNode->ninputs(), 4u);
298TEST(NodeHoistingTests, statefulOperations)
304 auto controlType = ControlType::Create(2);
305 auto valueType = TestType::createValueType();
306 auto stateType = TestType::createStateType();
307 const auto functionType = FunctionType::Create(
316 auto & rvsdg = rvsdgModule.Rvsdg();
318 auto lambdaNode = LambdaNode::Create(
319 rvsdg.GetRootRegion(),
321 auto controlArgument = lambdaNode->GetFunctionArguments()[0];
322 auto valueArgument = lambdaNode->GetFunctionArguments()[1];
323 auto stateArgument = lambdaNode->GetFunctionArguments()[2];
325 auto gammaNode1 = GammaNode::create(controlArgument, 2);
326 auto controlEntryVar = gammaNode1->AddEntryVar(controlArgument);
327 auto valueEntryVar1 = gammaNode1->AddEntryVar(valueArgument);
328 auto stateEntryVar = gammaNode1->AddEntryVar(stateArgument);
330 auto stateNode = TestOperation::createNode(
331 gammaNode1->subregion(0),
332 { valueEntryVar1.branchArgument[0], stateEntryVar.branchArgument[0] },
335 auto gammaNode2 = GammaNode::create(controlEntryVar.branchArgument[0], 2);
336 auto valueEntryVar2 = gammaNode2->AddEntryVar(stateNode->output(0));
337 auto valueEntryVar3 = gammaNode2->AddEntryVar(valueEntryVar1.branchArgument[0]);
339 auto binaryNode = TestOperation::createNode(
340 gammaNode2->subregion(0),
341 { valueEntryVar2.branchArgument[0], valueEntryVar3.branchArgument[0] },
345 gammaNode2->AddExitVar({ binaryNode->output(0), valueEntryVar2.branchArgument[1] });
347 auto exitVar = gammaNode1->AddExitVar({ exitVar2.output, valueEntryVar1.branchArgument[1] });
349 lambdaNode->finalize({ exitVar.output });
365 EXPECT_EQ(lambdaNode->subregion()->numNodes(), 1u);
368 EXPECT_EQ(gammaNode1->subregion(0)->numNodes(), 3u);
369 EXPECT_EQ(gammaNode1->subregion(1)->numNodes(), 0u);
371 EXPECT_EQ(gammaNode2->subregion(0)->numNodes(), 0u);
372 EXPECT_EQ(gammaNode2->subregion(1)->numNodes(), 0u);
375TEST(NodeHoistingTests, controlConstants)
410 auto controlType = ControlType::Create(2);
411 auto int32Type = BitType::Create(32);
414 const auto functionType = FunctionType::Create(
415 { ioStateType, memoryStateType },
416 { int32Type, ioStateType, memoryStateType });
419 auto & rvsdg = rvsdgModule.Rvsdg();
421 auto lambdaNode = LambdaNode::Create(
422 rvsdg.GetRootRegion(),
425 auto ioStateArgument = lambdaNode->GetFunctionArguments()[0];
426 auto memStateArgument = lambdaNode->GetFunctionArguments()[1];
429 auto thetaNode = ThetaNode::create(lambdaNode->subregion());
431 auto thetaCtrlLoopVar = thetaNode->AddLoopVar(thetaUndef);
432 auto & thetaInnerCtrl1 = ControlConstantOperation::create(*thetaNode->subregion(), 2, 1);
433 thetaCtrlLoopVar.post->divert_to(&thetaInnerCtrl1);
436 auto & gamma1 = GammaNode::Create(*thetaCtrlLoopVar.output, 2, {});
439 auto & gamma1Ctrl1 = ControlConstantOperation::create(*gamma1.subregion(0), 2, 1);
440 auto & gamma1Ctrl0 = ControlConstantOperation::create(*gamma1.subregion(1), 2, 0);
441 auto gamma1Exit = gamma1.AddExitVar({ &gamma1Ctrl1, &gamma1Ctrl0 });
444 auto & gamma2 = GammaNode::Create(*gamma1Exit.output, 2, {});
450 auto gamma2Exit = gamma2.AddExitVar({ &gamma2Int3, &gamma2Int7 });
452 lambdaNode->finalize({ gamma2Exit.output, ioStateArgument, memStateArgument });
469 EXPECT_EQ(lambdaNode->subregion()->numNodes(), 6u);
472 EXPECT_EQ(thetaNode->subregion()->numNodes(), 2u);
473 auto thetaPredicateOwner = TryGetOwnerNode<SimpleNode>(*thetaNode->predicate()->origin());
474 EXPECT_TRUE(thetaPredicateOwner);
475 EXPECT_EQ(thetaPredicateOwner->region(), thetaNode->subregion());
476 auto thetaPostOwner = TryGetOwnerNode<SimpleNode>(*thetaCtrlLoopVar.post->origin());
477 EXPECT_TRUE(thetaPostOwner);
478 EXPECT_EQ(thetaPostOwner->region(), thetaNode->subregion());
481 EXPECT_EQ(gamma1.subregion(0)->numNodes(), 1u);
482 auto gamma1LeftCtrlOwner = TryGetOwnerNode<SimpleNode>(*gamma1Exit.branchResult[0]->origin());
483 EXPECT_TRUE(gamma1LeftCtrlOwner);
484 EXPECT_EQ(gamma1LeftCtrlOwner->region(), gamma1.subregion(0));
486 EXPECT_EQ(gamma1.subregion(1)->numNodes(), 1u);
487 auto gamma1RightCtrlOwner = TryGetOwnerNode<SimpleNode>(*gamma1Exit.branchResult[1]->origin());
488 EXPECT_TRUE(gamma1RightCtrlOwner);
489 EXPECT_EQ(gamma1RightCtrlOwner->region(), gamma1.subregion(1));
492 EXPECT_EQ(gamma2.subregion(0)->numNodes(), 0u);
493 EXPECT_EQ(gamma2.subregion(1)->numNodes(), 0u);
496TEST(NodeHoistingTests, hoistLoadNodesOutOfGamma)
502 const auto i32Type = BitType::Create(32);
505 const auto controlType = ControlType::Create(2);
506 const auto functionType = FunctionType::Create(
507 { controlType, ptrType, ioStateType, memoryStateType },
508 { i32Type, ioStateType, memoryStateType });
511 auto & rvsdg = rvsdgModule.Rvsdg();
513 auto lambdaNode = LambdaNode::Create(
514 rvsdg.GetRootRegion(),
516 auto controlArgument = lambdaNode->GetFunctionArguments()[0];
517 auto ptrArgument = lambdaNode->GetFunctionArguments()[1];
518 auto ioStateArgument = lambdaNode->GetFunctionArguments()[2];
519 auto memoryStateArgument = lambdaNode->GetFunctionArguments()[3];
521 auto gammaNode = GammaNode::create(controlArgument, 2);
522 auto ptrEntryVar = gammaNode->AddEntryVar(ptrArgument);
523 auto ioStateEntryVar = gammaNode->AddEntryVar(ioStateArgument);
524 auto memoryStateEntryVar = gammaNode->AddEntryVar(memoryStateArgument);
528 *ptrEntryVar.branchArgument[0],
529 *ioStateEntryVar.branchArgument[0],
532 *hoistBarrierNode.output(0),
533 { memoryStateEntryVar.branchArgument[0] },
539 *ptrEntryVar.branchArgument[1],
540 { memoryStateEntryVar.branchArgument[1] },
544 *ptrEntryVar.branchArgument[1],
545 *loadNode1.output(0),
546 { loadNode1.output(1) },
549 auto i32ExitVar = gammaNode->AddExitVar({ loadNode0.output(0), loadNode1.output(0) });
550 auto ioStateExitVar = gammaNode->AddExitVar(
551 { ioStateEntryVar.branchArgument[0], ioStateEntryVar.branchArgument[1] });
552 auto memoryStateExitVar = gammaNode->AddExitVar({ loadNode0.output(1), storeNode1.output(0) });
555 lambdaNode->finalize({ i32ExitVar.output, ioStateExitVar.output, memoryStateExitVar.output });
557 GraphExport::Create(*lambdaOutput,
"x");
566 EXPECT_EQ(lambdaNode->subregion()->numNodes(), 2u);
567 EXPECT_EQ(gammaNode->subregion(0)->numNodes(), 2u);
568 EXPECT_EQ(gammaNode->subregion(1)->numNodes(), 1u);
573 EXPECT_EQ(gammaNode->ninputs(), 5u);
576TEST(NodeHoistingTests, hoistLoadNodesOutofNestedGamma)
582 const auto i32Type = BitType::Create(32);
585 const auto controlType = ControlType::Create(2);
586 const auto functionType = FunctionType::Create(
587 { controlType, ptrType, ptrType, ioStateType, memoryStateType },
588 { i32Type, ioStateType, memoryStateType });
591 auto & rvsdg = rvsdgModule.Rvsdg();
593 auto lambdaNode = LambdaNode::Create(
594 rvsdg.GetRootRegion(),
596 auto controlArgument = lambdaNode->GetFunctionArguments()[0];
597 auto ptrArgument1 = lambdaNode->GetFunctionArguments()[1];
598 auto ptrArgument2 = lambdaNode->GetFunctionArguments()[2];
599 auto ioStateArgument = lambdaNode->GetFunctionArguments()[3];
600 auto memoryStateArgument = lambdaNode->GetFunctionArguments()[4];
602 auto outerGammaNode = GammaNode::create(controlArgument, 2);
603 auto ctlEntryVar = outerGammaNode->AddEntryVar(controlArgument);
604 auto outerPtr1EntryVar = outerGammaNode->AddEntryVar(ptrArgument1);
605 auto outerPtr2EntryVar = outerGammaNode->AddEntryVar(ptrArgument2);
606 auto outerIOStateEntryVar = outerGammaNode->AddEntryVar(ioStateArgument);
607 auto outerMemoryStateEntryVar = outerGammaNode->AddEntryVar(memoryStateArgument);
611 *outerPtr1EntryVar.branchArgument[0],
612 *outerIOStateEntryVar.branchArgument[0],
615 auto innerGammaNode = GammaNode::create(ctlEntryVar.branchArgument[0], 2);
616 auto innerPtr1EntryVar = innerGammaNode->AddEntryVar(hoistBarrierNode.output(0));
617 auto innerPtr2EntryVar = innerGammaNode->AddEntryVar(outerPtr2EntryVar.branchArgument[0]);
618 auto innerMemoryStateEntryVar =
619 innerGammaNode->AddEntryVar(outerMemoryStateEntryVar.branchArgument[0]);
623 *innerPtr1EntryVar.branchArgument[0],
624 { innerMemoryStateEntryVar.branchArgument[0] },
630 *innerPtr2EntryVar.branchArgument[1],
631 { innerMemoryStateEntryVar.branchArgument[1] },
636 auto innerI32ExitVar = innerGammaNode->AddExitVar({ loadNode1.output(0), loadNode2.output(0) });
637 auto innerMemoryStateExitVar = innerGammaNode->AddExitVar(
638 { loadNode1.output(1), innerMemoryStateEntryVar.branchArgument[1] });
641 auto test1 = TestOperation::createNode(outerGammaNode->subregion(1), {}, { i32Type });
644 auto outerI32ExitVar = outerGammaNode->AddExitVar({ innerI32ExitVar.output, test1->output(0) });
645 auto outerIOStateExitVar = outerGammaNode->AddExitVar(
646 { outerIOStateEntryVar.branchArgument[0], outerIOStateEntryVar.branchArgument[1] });
647 auto outerMemoryStateExitVar = outerGammaNode->AddExitVar(
648 { innerMemoryStateExitVar.output, outerMemoryStateEntryVar.branchArgument[1] });
651 auto lambdaOutput = lambdaNode->finalize(
652 { outerI32ExitVar.output, outerIOStateExitVar.output, outerMemoryStateExitVar.output });
654 GraphExport::Create(*lambdaOutput,
"x");
672 Region::containsOperation<LoadNonVolatileOperation>(*innerGammaNode->subregion(0),
false));
674 Region::containsOperation<LoadNonVolatileOperation>(*innerGammaNode->subregion(1),
false));
677 Region::containsOperation<LoadNonVolatileOperation>(*outerGammaNode->subregion(0),
false));
678 EXPECT_EQ(outerGammaNode->subregion(0)->numNodes(), 3u);
680 EXPECT_TRUE(Region::containsOperation<LoadNonVolatileOperation>(*lambdaNode->subregion(),
false));
681 EXPECT_EQ(lambdaNode->subregion()->numNodes(), 3u);
684 auto [hoistedLoadNode1, loadOp] =
686 *innerMemoryStateEntryVar.input->origin());
687 EXPECT_NE(loadOp,
nullptr);
692 auto [hoistedLoadNode2, loadOp] =
694 *outerMemoryStateEntryVar.input->origin());
695 EXPECT_NE(loadOp,
nullptr);
700TEST(NodeHoistingTests, hoistLoadNodeOutOfGammaInTheta)
706 const auto i32Type = BitType::Create(32);
708 const auto controlType = ControlType::Create(2);
709 const auto functionType =
710 FunctionType::Create({ controlType, ptrType, memoryStateType }, { ptrType, memoryStateType });
713 auto & rvsdg = rvsdgModule.Rvsdg();
715 auto lambdaNode = LambdaNode::Create(
716 rvsdg.GetRootRegion(),
718 auto controlArgument = lambdaNode->GetFunctionArguments()[0];
719 auto ptrArgument = lambdaNode->GetFunctionArguments()[1];
720 auto memoryStateArgument = lambdaNode->GetFunctionArguments()[2];
722 auto thetaNode = ThetaNode::create(lambdaNode->subregion());
723 auto ptrLoopVar = thetaNode->AddLoopVar(ptrArgument);
724 auto memoryStateLoopVar = thetaNode->AddLoopVar(memoryStateArgument);
725 auto controlLoopVar = thetaNode->AddLoopVar(controlArgument);
727 auto gammaNode = GammaNode::create(controlLoopVar.pre, 2);
728 auto ptrEntryVar = gammaNode->AddEntryVar(ptrLoopVar.pre);
729 auto memoryStateEntryVar = gammaNode->AddEntryVar(memoryStateLoopVar.pre);
732 auto testNode = TestOperation::createNode(gammaNode->subregion(0), {}, { i32Type });
736 *ptrEntryVar.branchArgument[1],
737 { memoryStateEntryVar.branchArgument[1] },
741 auto i32ExitVar = gammaNode->AddExitVar({ testNode->output(0), loadNode.output(0) });
742 auto memoryStateExitVar =
743 gammaNode->AddExitVar({ memoryStateEntryVar.branchArgument[0], loadNode.output(1) });
745 memoryStateLoopVar.post->divert_to(memoryStateExitVar.output);
747 auto lambdaOutput = lambdaNode->finalize({ ptrLoopVar.output, memoryStateLoopVar.output });
749 GraphExport::Create(*lambdaOutput,
"x");
758 EXPECT_TRUE(Region::containsOperation<LoadNonVolatileOperation>(*thetaNode->subregion(),
false));
762 EXPECT_EQ(gammaNode->ninputs(), 5u);