Jlm
Loading...
Searching...
No Matches
GammaConversionTests.cpp
Go to the documentation of this file.
1/*
2 * Copyright 2022 Nico Reißmann <nico.reissmann@gmail.com>
3 * Magnus Sjalander <work@sjalander.com>
4 * See COPYING for terms of redistribution.
5 */
6
7#include <gtest/gtest.h>
8
10#include <jlm/hls/ir/hls.hpp>
14#include <jlm/rvsdg/gamma.hpp>
15#include <jlm/rvsdg/lambda.hpp>
18
19namespace jlm::hls
20{
21
22using namespace jlm::rvsdg;
23using namespace jlm::llvm;
24using namespace jlm::util;
25
26static void
28{
31
32 auto & rootRegion = rvsdgModule.Rvsdg().GetRootRegion();
33
34 JLM_ASSERT(rootRegion.numNodes() == 1);
35 auto * lambda = dynamic_cast<LambdaNode *>(&*rootRegion.Nodes().begin());
36 JLM_ASSERT(lambda != nullptr);
37 EXPECT_FALSE(Region::containsOperation<GammaOperation>(rootRegion, true));
38
39 size_t muxCount = 0;
40 for (auto & subnode : TopDownTraverser(lambda->subregion()))
41 {
42 if (is<MuxOperation>(subnode->GetOperation()))
43 {
44 muxCount++;
45 }
46 }
48}
49
51{
53 auto functionType =
55
57
58 auto lambda = LambdaNode::Create(
59 rvsdgModule.Rvsdg().GetRootRegion(),
60 LlvmLambdaOperation::Create(functionType, "f", Linkage::externalLinkage));
61
62 auto & matchNode =
63 MatchOperation::CreateNode(*lambda->GetFunctionArguments()[0], { { 0, 0 } }, 1, 2);
64 auto gamma = GammaNode::create(matchNode.output(0), 2);
65 auto entryVar1 = gamma->AddEntryVar(lambda->GetFunctionArguments()[1]);
66 auto entryVar2 = gamma->AddEntryVar(lambda->GetFunctionArguments()[2]);
67 auto exitVar = gamma->AddExitVar({ entryVar1.branchArgument[0], entryVar2.branchArgument[1] });
68
69 auto lambdaOutput = lambda->finalize({ exitVar.output });
71
73}
74
75TEST(GammaConversionTests, WithoutMatchOperation)
76{
77 auto valueType = TestType::createValueType();
78 auto functionType =
79 FunctionType::Create({ ControlType::Create(2), valueType, valueType }, { valueType });
80
81 LlvmRvsdgModule rvsdgModule(FilePath(""), "", "");
82
83 auto lambda = LambdaNode::Create(
84 rvsdgModule.Rvsdg().GetRootRegion(),
85 LlvmLambdaOperation::Create(functionType, "f", Linkage::externalLinkage));
86
87 auto gamma = GammaNode::create(lambda->GetFunctionArguments()[0], 2);
88 auto entryVar1 = gamma->AddEntryVar(lambda->GetFunctionArguments()[1]);
89 auto entryVar2 = gamma->AddEntryVar(lambda->GetFunctionArguments()[2]);
90 auto exitVar = gamma->AddExitVar({ entryVar1.branchArgument[0], entryVar2.branchArgument[1] });
91
92 auto lambdaOutput = lambda->finalize({ exitVar.output });
93 GraphExport::Create(*lambdaOutput, "");
94
95 TestGammaConversion(rvsdgModule, 1);
96}
97
98TEST(GammaConversionTests, NestedGammas)
99{
100 auto controlType = ControlType::Create(2);
101 auto bit32Type = BitType::Create(32);
102 auto functionType = FunctionType::Create({ controlType }, { bit32Type });
103
104 LlvmRvsdgModule rvsdgModule(FilePath(""), "", "");
105
106 auto lambda = LambdaNode::Create(
107 rvsdgModule.Rvsdg().GetRootRegion(),
108 LlvmLambdaOperation::Create(functionType, "f", Linkage::externalLinkage));
109
110 auto & outerControlConstant =
111 BitConstantOperation::create(*lambda->subregion(), BitValueRepresentation(32, 0));
112 auto outerGamma = GammaNode::create(lambda->GetFunctionArguments()[0], 2);
113 auto outerVar = outerGamma->AddEntryVar(&outerControlConstant);
114
115 auto & innerControlConstant =
116 ControlConstantOperation::create(*outerGamma->subregion(1), ControlValueRepresentation(0, 2));
117 auto innerGamma = GammaNode::create(&innerControlConstant, 2);
118 auto innerVar = innerGamma->AddEntryVar(outerVar.branchArgument[1]);
119 auto innerExit =
120 innerGamma->AddExitVar({ innerVar.branchArgument[0], innerVar.branchArgument[1] });
121
122 auto outerExit = outerGamma->AddExitVar({ outerVar.branchArgument[0], innerExit.output });
123
124 auto lambdaOutput = lambda->finalize({ outerExit.output });
125 GraphExport::Create(*lambdaOutput, "");
126
127 TestGammaConversion(rvsdgModule, 2);
128}
129
130TEST(GammaConversionTests, MuxPredicateMapping)
131{
132 auto valueType = TestType::createValueType();
133 auto bitType = BitType::Create(1);
134 auto functionType = FunctionType::Create({ bitType, valueType, valueType }, { valueType });
135
136 LlvmRvsdgModule rvsdgModule(FilePath(""), "", "");
137
138 auto lambda = LambdaNode::Create(
139 rvsdgModule.Rvsdg().GetRootRegion(),
140 LlvmLambdaOperation::Create(functionType, "f", Linkage::externalLinkage));
141
142 auto & matchNode =
143 MatchOperation::CreateNode(*lambda->GetFunctionArguments()[0], { { 0, 0 } }, 1, 2);
144
145 auto gamma = GammaNode::create(matchNode.output(0), 2);
146 auto entryVar1 = gamma->AddEntryVar(lambda->GetFunctionArguments()[1]);
147 auto entryVar2 = gamma->AddEntryVar(lambda->GetFunctionArguments()[2]);
148 auto exitVar = gamma->AddExitVar({ entryVar1.branchArgument[0], entryVar2.branchArgument[1] });
149
150 auto lambdaOutput = lambda->finalize({ exitVar.output });
151 GraphExport::Create(*lambdaOutput, "");
152
153 TestGammaConversion(rvsdgModule, 1);
154
155 for (auto & node : TopDownTraverser(lambda->subregion()))
156 {
157 if (is<MuxOperation>(node->GetOperation()))
158 {
159 EXPECT_EQ(node->input(0)->origin(), matchNode.output(0));
160 }
161 }
162}
163
164TEST(GammaConversionTests, MuxAlternativeSelection)
165{
166 auto valueType = TestType::createValueType();
167 auto functionType = FunctionType::Create({ BitType::Create(2), valueType }, { valueType });
168
169 LlvmRvsdgModule rvsdgModule(FilePath(""), "", "");
170
171 auto lambda = LambdaNode::Create(
172 rvsdgModule.Rvsdg().GetRootRegion(),
173 LlvmLambdaOperation::Create(functionType, "f", Linkage::externalLinkage));
174
175 auto & matchNode = MatchOperation::CreateNode(
176 *lambda->GetFunctionArguments()[0],
177 { { 0, 0 }, { 1, 1 }, { 2, 2 } },
178 3,
179 3);
180
181 auto gamma = GammaNode::create(matchNode.output(0), 3);
182 auto entryVar1 = gamma->AddEntryVar(lambda->GetFunctionArguments()[1]);
183 auto exitVar = gamma->AddExitVar(
184 { entryVar1.branchArgument[0], entryVar1.branchArgument[1], entryVar1.branchArgument[2] });
185
186 auto lambdaOutput = lambda->finalize({ exitVar.output });
187 GraphExport::Create(*lambdaOutput, "");
188
189 TestGammaConversion(rvsdgModule, 1);
190
191 for (auto & node : TopDownTraverser(lambda->subregion()))
192 {
193 if (is<MuxOperation>(node->GetOperation()))
194 {
195 auto & muxOp = static_cast<const MuxOperation &>(node->GetOperation());
196 EXPECT_EQ(muxOp.narguments(), 4u);
197 EXPECT_EQ(node->ninputs(), 4u);
198 }
199 }
200}
201
202TEST(GammaConversionTests, MuxControlPredicateMapping)
203{
204 auto valueType = TestType::createValueType();
205 auto functionType =
206 FunctionType::Create({ BitType::Create(1), valueType, valueType }, { valueType });
207
208 LlvmRvsdgModule rvsdgModule(FilePath(""), "", "");
209
210 auto lambda = LambdaNode::Create(
211 rvsdgModule.Rvsdg().GetRootRegion(),
212 LlvmLambdaOperation::Create(functionType, "f", Linkage::externalLinkage));
213
214 auto & matchNode =
215 MatchOperation::CreateNode(*lambda->GetFunctionArguments()[0], { { 0, 0 } }, 1, 2);
216
217 auto gamma = GammaNode::create(matchNode.output(0), 2);
218 auto entryVar1 = gamma->AddEntryVar(lambda->GetFunctionArguments()[1]);
219 auto entryVar2 = gamma->AddEntryVar(lambda->GetFunctionArguments()[2]);
220 auto exitVar = gamma->AddExitVar({ entryVar1.branchArgument[0], entryVar2.branchArgument[1] });
221
222 auto lambdaOutput = lambda->finalize({ exitVar.output });
223 GraphExport::Create(*lambdaOutput, "");
224
225 TestGammaConversion(rvsdgModule, 1);
226
227 for (auto & node : TopDownTraverser(lambda->subregion()))
228 {
229 if (is<MuxOperation>(node->GetOperation()))
230 {
231 EXPECT_EQ(node->ninputs(), 3u);
232
233 auto * muxPredicate = node->input(0)->origin();
234 EXPECT_EQ(muxPredicate, matchNode.output(0));
235 }
236 }
237}
238
239TEST(GammaConversionTests, SpeculativeConversionUsesDiscardingMux)
240{
241 auto valueType = TestType::createValueType();
242 auto functionType = FunctionType::Create({ BitType::Create(2), valueType }, { valueType });
243
244 LlvmRvsdgModule rvsdgModule(FilePath(""), "", "");
245
246 auto lambda = LambdaNode::Create(
247 rvsdgModule.Rvsdg().GetRootRegion(),
248 LlvmLambdaOperation::Create(functionType, "f", Linkage::externalLinkage));
249
250 auto & matchNode = MatchOperation::CreateNode(
251 *lambda->GetFunctionArguments()[0],
252 { { 0, 0 }, { 1, 1 }, { 2, 2 } },
253 3,
254 3);
255
256 auto gamma = GammaNode::create(matchNode.output(0), 3);
257 auto entryVar1 = gamma->AddEntryVar(lambda->GetFunctionArguments()[1]);
258 auto exitVar = gamma->AddExitVar(
259 { entryVar1.branchArgument[0], entryVar1.branchArgument[1], entryVar1.branchArgument[2] });
260
261 auto lambdaOutput = lambda->finalize({ exitVar.output });
262 GraphExport::Create(*lambdaOutput, "");
263
264 TestGammaConversion(rvsdgModule, 1);
265
266 for (auto & node : TopDownTraverser(lambda->subregion()))
267 {
268 if (is<MuxOperation>(node->GetOperation()))
269 {
270 auto & muxOp = static_cast<const MuxOperation &>(node->GetOperation());
271 EXPECT_TRUE(muxOp.discarding);
272 }
273 }
274}
275
276TEST(GammaConversionTests, NonSpeculativeModeUsesBranches)
277{
278 auto valueType = TestType::createValueType();
279 auto stateType = TestType::createStateType();
280 auto functionType =
281 FunctionType::Create({ BitType::Create(2), valueType, valueType, stateType }, { stateType });
282
283 LlvmRvsdgModule rvsdgModule(FilePath(""), "", "");
284
285 auto lambda = LambdaNode::Create(
286 rvsdgModule.Rvsdg().GetRootRegion(),
287 LlvmLambdaOperation::Create(functionType, "f", Linkage::externalLinkage));
288
289 auto & matchNode = MatchOperation::CreateNode(
290 *lambda->GetFunctionArguments()[0],
291 { { 0, 0 }, { 1, 1 }, { 2, 2 } },
292 3,
293 3);
294
295 auto gamma = GammaNode::create(matchNode.output(0), 3);
296 auto stateVar = gamma->AddEntryVar(lambda->GetFunctionArguments()[3]);
297 auto stateExit = gamma->AddExitVar(
298 { stateVar.branchArgument[0], stateVar.branchArgument[1], stateVar.branchArgument[2] });
299
300 auto lambdaOutput = lambda->finalize({ stateExit.output });
301 GraphExport::Create(*lambdaOutput, "");
302
303 TestGammaConversion(rvsdgModule, 1);
304
305 size_t branchCount = 0;
306 for (auto & node : TopDownTraverser(lambda->subregion()))
307 {
308 if (is<BranchOperation>(node->GetOperation()))
309 {
310 EXPECT_EQ(node->ninputs(), 2u);
311 branchCount++;
312 }
313 }
314 EXPECT_GE(branchCount, 1u);
315}
316
317} // namespace
TEST(BaseHlsTests, TestIsForbiddenChar)
static jlm::util::StatisticsCollector statisticsCollector
static void CreateAndRun(rvsdg::RvsdgModule &rvsdgModule, util::StatisticsCollector &statisticsCollector)
static std::unique_ptr< LlvmLambdaOperation > Create(std::shared_ptr< const jlm::rvsdg::FunctionType > type, std::string name, const jlm::llvm::Linkage &linkage, jlm::llvm::CallingConvention callingConvention, jlm::llvm::AttributeSet attributes)
Definition lambda.hpp:84
static std::shared_ptr< const BitType > Create(std::size_t nbits)
Creates bit type of specified width.
Definition type.cpp:45
static std::shared_ptr< const FunctionType > Create(std::vector< std::shared_ptr< const jlm::rvsdg::Type > > argumentTypes, std::vector< std::shared_ptr< const jlm::rvsdg::Type > > resultTypes)
static GammaNode * create(jlm::rvsdg::Output *predicate, size_t nalternatives)
Definition gamma.hpp:161
static GraphExport & Create(Output &origin, std::string name)
Definition graph.cpp:62
static LambdaNode * Create(rvsdg::Region &parent, std::unique_ptr< LambdaOperation > operation)
Definition lambda.cpp:141
static Node & CreateNode(Output &predicate, const std::unordered_map< uint64_t, uint64_t > &mapping, const uint64_t defaultAlternative, const size_t numAlternatives)
Definition control.hpp:220
static std::shared_ptr< const TestType > createValueType()
Definition TestType.cpp:67
#define JLM_ASSERT(x)
Definition common.hpp:16
static void TestGammaConversion(RvsdgModule &rvsdgModule, size_t expectedMuxCount)
TEST(GammaConversionTests, WithMatchOperation)
Global memory state passed between functions.
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.