Jlm
Loading...
Searching...
No Matches
hls-function-util.cpp
Go to the documentation of this file.
1//
2// Created by david on 7/2/21.
3//
4
6#include <jlm/hls/ir/hls.hpp>
12#include <jlm/rvsdg/gamma.hpp>
16#include <jlm/rvsdg/theta.hpp>
18#include <jlm/rvsdg/view.hpp>
19
20#include <deque>
21
22namespace jlm::hls
23{
24
25std::vector<rvsdg::LambdaNode::ContextVar>
27{
28 std::vector<rvsdg::LambdaNode::ContextVar> result;
29 for (auto cv : lambda->GetContextVars())
30 {
31 auto ip = cv.input;
34 auto arg = util::assertedCast<const llvm::LlvmGraphImport>(traced);
35 if (dynamic_cast<const rvsdg::FunctionType *>(arg->ImportedType().get())
36 && arg->Name().find(name_contains) != arg->Name().npos)
37 {
38 result.push_back(cv);
39 }
40 }
41 return result;
42}
43
44void
46 rvsdg::Output * output,
47 std::vector<rvsdg::SimpleNode *> & calls,
48 std::unordered_set<rvsdg::Output *> & visited)
49{
50 if (visited.count(output))
51 {
52 // skip already processed outputs
53 return;
54 }
55 visited.insert(output);
56 for (auto & user : output->Users())
57 {
59 user.GetOwner(),
60 [&](rvsdg::Node * node)
61 {
62 rvsdg::MatchTypeOrFail(
63 *node,
64 [&](rvsdg::SimpleNode & simplenode)
65 {
66 rvsdg::MatchTypeWithDefault(
67 simplenode.GetOperation(),
68 [&](const llvm::CallOperation &)
69 {
70 // TODO: verify this is the right type of function call
71 calls.push_back(&simplenode);
72 },
73 [&]()
74 {
75 for (size_t i = 0; i < simplenode.noutputs(); ++i)
76 {
77 trace_function_calls(simplenode.output(i), calls, visited);
78 }
79 });
80 },
81 [&](LoopNode & loop)
82 {
83 trace_function_calls(loop.mapInput(user).inner, calls, visited);
84 },
85 [&](rvsdg::ThetaNode & theta)
86 {
87 trace_function_calls(theta.MapInputLoopVar(user).pre, calls, visited);
88 },
89 [&](rvsdg::GammaNode & gamma)
90 {
91 rvsdg::MatchVariant(
92 gamma.MapInput(user),
93 [&](const rvsdg::GammaNode::MatchVar &)
94 {
95 },
97 {
98 for (auto out : evar.branchArgument)
99 {
100 trace_function_calls(out, calls, visited);
101 }
102 });
103 });
104 },
105 [&](rvsdg::Region * region)
106 {
107 rvsdg::MatchTypeOrFail(
108 *region->node(),
109 [&](LoopNode & loop)
110 {
111 rvsdg::MatchVariant(
112 loop.mapResult(user),
113 [&](const LoopNode::BackEdgeVar & backedge)
114 {
115 trace_function_calls(backedge.pre, calls, visited);
116 },
117 [&](const LoopNode::ExitVar & exit)
118 {
119 trace_function_calls(exit.output, calls, visited);
120 });
121 },
122 [&](rvsdg::ThetaNode & theta)
123 {
124 rvsdg::MatchVariant(
125 theta.mapResult(user),
126 [&](const rvsdg::ThetaNode::LoopVar & loopvar)
127 {
128 trace_function_calls(loopvar.output, calls, visited);
129 },
130 [&](const rvsdg::ThetaNode::PredicateVar &)
131 {
132 });
133 },
134 [&](rvsdg::GammaNode & gamma)
135 {
136 trace_function_calls(gamma.MapBranchResultExitVar(user).output, calls, visited);
137 });
138 });
139 }
140}
141
142const llvm::IntegerConstantOperation *
144{
145 if (auto arg = dynamic_cast<const rvsdg::RegionArgument *>(dst))
146 {
147 return trace_constant(arg->input()->origin());
148 }
149
150 auto [constantNode, constantOperation] =
151 rvsdg::TryGetSimpleNodeAndOptionalOp<llvm::IntegerConstantOperation>(*dst);
152 if (constantNode)
153 {
154 if (constantOperation)
155 return constantOperation;
156
157 for (size_t i = 0; i < constantNode->ninputs(); ++i)
158 {
159 // TODO: fix, this is a hack - only works because of distribute constants
160 if (*constantNode->input(i)->Type() == *dst->Type())
161 {
162 return trace_constant(constantNode->input(i)->origin());
163 }
164 }
165 }
166
167 JLM_UNREACHABLE("Constant not found");
168}
169
172{
173 // create lists of nested regions
174 std::deque<rvsdg::Region *> target_regions = get_parent_regions(target);
175 std::deque<rvsdg::Region *> out_regions = get_parent_regions(out->region());
176 JLM_ASSERT(target_regions.front() == out_regions.front());
177 // remove common ancestor regions
178 rvsdg::Region * common_region = nullptr;
179 while (!target_regions.empty() && !out_regions.empty()
180 && target_regions.front() == out_regions.front())
181 {
182 common_region = target_regions.front();
183 target_regions.pop_front();
184 out_regions.pop_front();
185 }
186 // route out to convergence point from out
187 rvsdg::Output * common_out = route_request_rhls(common_region, out);
188 auto common_loop = dynamic_cast<LoopNode *>(common_region->node());
189 if (common_loop)
190 {
191 // add a backedge to prevent cycles
192 auto arg = common_loop->add_backedge(out->Type());
193 arg->result()->divert_to(common_out);
194 // route inwards from convergence point to target
195 auto result = route_response_rhls(target, arg);
196 return result;
197 }
198 else
199 {
200 // lambda is common region - might create cycle
201 // TODO: how to check that this won't create a cycle
203 target_regions.empty() || target_regions.front()->node()->region() == common_out->region());
204 return route_response_rhls(target, common_out);
205 }
206}
207
210{
211 if (response->region() == target)
212 {
213 return response;
214 }
215 else
216 {
217 auto parent_response = route_response_rhls(target->node()->region(), response);
218 auto ln = util::assertedCast<LoopNode>(target->node());
219 return ln->addResponseInput(parent_response);
220 }
221}
222
225{
226 if (request->region() == target)
227 {
228 return request;
229 }
230
231 auto ln = util::assertedCast<LoopNode>(request->region()->node());
232 auto output = ln->addRequestOutput(request);
233
234 return route_request_rhls(target, output);
235}
236
237std::deque<rvsdg::Region *>
239{
240 std::deque<rvsdg::Region *> regions;
241 rvsdg::Region * target_region = region;
242 while (!dynamic_cast<const llvm::LlvmLambdaOperation *>(&target_region->node()->GetOperation()))
243 {
244 regions.push_front(target_region);
245 target_region = target_region->node()->region();
246 }
247 regions.push_front(target_region);
248 return regions;
249}
250
251const rvsdg::Output *
253{
254 return rvsdg::MatchVariant(
255 output->GetOwner(),
256 [&](rvsdg::Region * region) -> const rvsdg::Output *
257 {
258 if (region->IsRootRegion())
259 {
260 return output;
261 }
262 return rvsdg::MatchTypeOrFail(
263 *region->node(),
264 [&](LoopNode & loop) -> const rvsdg::Output *
265 {
266 (void)loop;
267 if (dynamic_cast<const BackEdgeArgument *>(output))
268 {
269 // don't follow backedges to avoid cycles
270 return nullptr;
271 }
272 return trace_call_rhls(dynamic_cast<const rvsdg::RegionArgument *>(output)->input());
273 },
274 [&](rvsdg::StructuralNode & structural) -> const rvsdg::Output *
275 {
276 (void)structural;
277 if (dynamic_cast<const BackEdgeArgument *>(output))
278 {
279 // don't follow backedges to avoid cycles
280 return nullptr;
281 }
282 return trace_call_rhls(dynamic_cast<const rvsdg::RegionArgument *>(output)->input());
283 });
284 },
285 [&](rvsdg::Node * node) -> const rvsdg::Output *
286 {
287 return rvsdg::MatchTypeOrFail(
288 *node,
289 [&](LoopNode & loop) -> const rvsdg::Output *
290 {
291 (void)loop;
292 auto so = dynamic_cast<const rvsdg::StructuralOutput *>(output);
293 for (auto & r : so->results)
294 {
295 if (auto result = trace_call_rhls(&r))
296 {
297 return result;
298 }
299 }
300 return nullptr;
301 },
302 [&](rvsdg::StructuralNode & structural) -> const rvsdg::Output *
303 {
304 (void)structural;
305 auto so = dynamic_cast<const rvsdg::StructuralOutput *>(output);
306 for (auto & r : so->results)
307 {
308 if (auto result = trace_call_rhls(&r))
309 {
310 return result;
311 }
312 }
313 return nullptr;
314 },
315 [&](rvsdg::SimpleNode & simple) -> const rvsdg::Output *
316 {
317 for (auto & input : simple.Inputs())
318 {
319 if (*input.Type() == *output->Type())
320 {
321 if (auto result = trace_call_rhls(&input))
322 {
323 return result;
324 }
325 }
326 }
327 return nullptr;
328 });
329 });
330}
331
332const rvsdg::Output *
334{
335 // version of trace call for rhls
336 return trace_call_rhls(input->origin());
337}
338
339bool
341{
342 auto ip = cv.input;
343 auto traced = trace_call_rhls(ip);
344 JLM_ASSERT(traced);
345 auto arg = util::assertedCast<const llvm::LlvmGraphImport>(traced);
346 return dynamic_cast<const rvsdg::FunctionType *>(arg->ImportedType().get());
347}
348
349std::string
351{
352 auto traced = jlm::hls::trace_call_rhls(input);
353 JLM_ASSERT(traced);
354 auto arg = jlm::util::assertedCast<const jlm::llvm::LlvmGraphImport>(traced);
355 return arg->Name();
356}
357
358bool
360{
361 if (dynamic_cast<const llvm::CallOperation *>(&node->GetOperation()))
362 {
363 auto name = get_function_name(node->input(0));
364 if (name.rfind("decouple_req") != name.npos)
365 return true;
366 }
367 return false;
368}
369
370bool
372{
373 if (dynamic_cast<const llvm::CallOperation *>(&node->GetOperation()))
374 {
375 auto name = get_function_name(node->input(0));
376 if (name.rfind("decouple_res") != name.npos)
377 return true;
378 }
379 return false;
380}
381
384{
385 JLM_ASSERT(state_edge);
386 JLM_ASSERT(state_edge->nusers() == 1);
387 JLM_ASSERT(rvsdg::is<llvm::MemoryStateType>(state_edge->Type()));
388 return &state_edge->SingleUser();
389}
390
393{
394 if (auto ba = dynamic_cast<BackEdgeArgument *>(out))
395 {
396 return FindSourceNode(ba->result()->origin());
397 }
398 else if (auto ra = dynamic_cast<rvsdg::RegionArgument *>(out))
399 {
400 if (ra->input() && rvsdg::TryGetOwnerNode<LoopNode>(*ra->input()))
401 {
402 return FindSourceNode(ra->input()->origin());
403 }
404 else
405 {
406 // lambda argument
407 return ra;
408 }
409 }
410 else if (auto so = dynamic_cast<rvsdg::StructuralOutput *>(out))
411 {
412 JLM_ASSERT(rvsdg::TryGetOwnerNode<LoopNode>(*out));
413 return FindSourceNode(so->results.begin()->origin());
414 }
415
416 JLM_ASSERT(rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(*out));
417 return out;
418}
419}
BackEdgeResult * result()
Definition hls.hpp:632
BackEdgeArgument * add_backedge(std::shared_ptr< const jlm::rvsdg::Type > type)
Definition hls.cpp:400
Call operation class.
Definition call.hpp:251
Function type class.
Conditional operator / pattern matching.
Definition gamma.hpp:99
void divert_to(Output *new_origin)
Definition node.cpp:64
Output * origin() const noexcept
Definition node.hpp:58
std::vector< ContextVar > GetContextVars() const noexcept
Gets all bound context variables.
Definition lambda.cpp:120
virtual const Operation & GetOperation() const noexcept=0
rvsdg::Region * region() const noexcept
Definition node.hpp:761
rvsdg::Input & SingleUser() noexcept
Definition node.hpp:347
rvsdg::Region * region() const noexcept
Definition node.cpp:151
UsersRange Users()
Definition node.hpp:354
std::variant< Node *, Region * > GetOwner() const noexcept
Definition node.hpp:378
const std::shared_ptr< const rvsdg::Type > & Type() const noexcept
Definition node.hpp:366
size_t nusers() const noexcept
Definition node.hpp:280
Represents the argument of a region.
Definition region.hpp:41
Represent acyclic RVSDG subgraphs.
Definition region.hpp:213
rvsdg::StructuralNode * node() const noexcept
Definition region.hpp:301
const SimpleOperation & GetOperation() const noexcept override
NodeInput * input(size_t index) const noexcept
#define JLM_ASSERT(x)
Definition common.hpp:16
#define JLM_UNREACHABLE(msg)
Definition common.hpp:43
rvsdg::Output * route_response_rhls(rvsdg::Region *target, rvsdg::Output *response)
void trace_function_calls(rvsdg::Output *output, std::vector< rvsdg::SimpleNode * > &calls, std::unordered_set< rvsdg::Output * > &visited)
std::deque< rvsdg::Region * > get_parent_regions(rvsdg::Region *region)
rvsdg::Output * FindSourceNode(rvsdg::Output *out)
bool is_function_argument(const rvsdg::LambdaNode::ContextVar &cv)
rvsdg::Output * route_request_rhls(rvsdg::Region *target, rvsdg::Output *request)
bool is_dec_res(rvsdg::SimpleNode *node)
std::string get_function_name(jlm::rvsdg::Input *input)
rvsdg::Output * route_to_region_rhls(rvsdg::Region *target, rvsdg::Output *out)
const llvm::IntegerConstantOperation * trace_constant(const rvsdg::Output *dst)
const rvsdg::Output * trace_call_rhls(const rvsdg::Output *output)
rvsdg::Input * get_mem_state_user(rvsdg::Output *state_edge)
std::vector< rvsdg::LambdaNode::ContextVar > find_function_arguments(const rvsdg::LambdaNode *lambda, std::string name_contains)
bool is_dec_req(rvsdg::SimpleNode *node)
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
A variable routed into all gamma regions.
Definition gamma.hpp:131
Bound context variable.
Definition lambda.hpp:100
rvsdg::Input * input
Input variable bound into lambda node.
Definition lambda.hpp:108