Jlm
Loading...
Searching...
No Matches
alloca-conv.cpp
Go to the documentation of this file.
1/*
2 * Copyright 2021 David Metz <david.c.metz@ntnu.no>
3 * See COPYING for terms of redistribution.
4 */
5
9#include <jlm/hls/ir/hls.hpp>
19#include <jlm/rvsdg/gamma.hpp>
23#include <jlm/rvsdg/theta.hpp>
25
26namespace jlm::hls
27{
28
30{
31public:
32 std::vector<jlm::rvsdg::SimpleNode *> load_nodes;
33 std::vector<jlm::rvsdg::SimpleNode *> store_nodes;
34
39
40private:
41 void
43 {
44 if (!rvsdg::is<llvm::PointerType>(op->Type()))
45 {
46 // only process pointer outputs
47 return;
48 }
49 if (visited.count(op))
50 {
51 // skip already processed outputs
52 return;
53 }
54 visited.insert(op);
55 for (auto & user : op->Users())
56 {
58 user.GetOwner(),
59 [&](rvsdg::Node * node)
60 {
61 rvsdg::MatchTypeOrFail(
62 *node,
63 [&](rvsdg::SimpleNode & simplenode)
64 {
65 rvsdg::MatchTypeWithDefault(
66 simplenode.GetOperation(),
67 [&](const jlm::llvm::StoreNonVolatileOperation &)
68 {
69 store_nodes.push_back(&simplenode);
70 },
71 [&](const jlm::llvm::LoadNonVolatileOperation &)
72 {
73 load_nodes.push_back(&simplenode);
74 },
75 [&](const jlm::llvm::CallOperation &)
76 {
77 // TODO: verify this is the right type of function call
78 throw util::Error("encountered a call for an alloca");
79 },
80 [&]()
81 {
82 for (size_t i = 0; i < simplenode.noutputs(); ++i)
83 {
84 trace(simplenode.output(i));
85 }
86 });
87 },
88 [&](LoopNode & loop)
89 {
90 trace(loop.mapInput(user).inner);
91 },
92 [&](rvsdg::ThetaNode & theta)
93 {
94 trace(theta.MapInputLoopVar(user).pre);
95 },
96 [&](rvsdg::GammaNode & gamma)
97 {
98 MatchVariant(
99 gamma.MapInput(user),
100 [&](const rvsdg::GammaNode::MatchVar &)
101 {
102 },
103 [&](const rvsdg::GammaNode::EntryVar & entry)
104 {
105 for (auto var : entry.branchArgument)
106 {
107 trace(var);
108 }
109 });
110 });
111 },
112 [&](rvsdg::Region * region)
113 {
115 *region->node(),
116 [&](LoopNode & loop)
117 {
118 rvsdg::MatchVariant(
119 loop.mapResult(user),
120 [&](const LoopNode::BackEdgeVar & backedge)
121 {
122 trace(backedge.pre);
123 },
124 [&](const LoopNode::ExitVar & exit)
125 {
126 trace(exit.output);
127 });
128 },
129 [&](rvsdg::ThetaNode & theta)
130 {
132 theta.mapResult(user),
133 [&](const rvsdg::ThetaNode::PredicateVar &)
134 {
135 },
136 [&](const rvsdg::ThetaNode::LoopVar & loopvar)
137 {
138 trace(loopvar.output);
139 });
140 },
141 [&](rvsdg::GammaNode & gamma)
142 {
143 trace(gamma.MapBranchResultExitVar(user).output);
144 });
145 });
146 }
147 }
148
149 std::unordered_set<jlm::rvsdg::Output *> visited;
150};
151
152static jlm::rvsdg::Output *
154{
155 // TODO: handle geps that are not direct predecessors
156 auto & node = rvsdg::AssertGetOwnerNode<rvsdg::SimpleNode>(*o);
157 util::assertedCast<const jlm::llvm::GetElementPtrOperation>(&node.GetOperation());
158 // pointer to array, i.e. first index is zero
159 // TODO: check
160 JLM_ASSERT(node.ninputs() == 3);
161 return node.input(2)->origin();
162}
163
164static void
166{
167 for (auto & node : rvsdg::TopDownTraverser(region))
168 {
169 if (auto structnode = dynamic_cast<rvsdg::StructuralNode *>(node))
170 {
171 for (size_t n = 0; n < structnode->nsubregions(); n++)
172 {
173 alloca_conv(structnode->subregion(n));
174 }
175 }
176 else if (auto po = dynamic_cast<const jlm::llvm::AllocaOperation *>(&(node->GetOperation())))
177 {
178 // ensure that the size is one
179 JLM_ASSERT(node->ninputs() == 1);
180 auto constant_output = dynamic_cast<rvsdg::NodeOutput *>(node->input(0)->origin());
181 JLM_ASSERT(constant_output);
182 auto constant_operation = dynamic_cast<const llvm::IntegerConstantOperation *>(
183 &constant_output->node()->GetOperation());
184 JLM_ASSERT(constant_operation);
185 JLM_ASSERT(constant_operation->Representation().to_uint() == 1);
186 // ensure that the alloca is an array type
187 auto at = std::dynamic_pointer_cast<const llvm::ArrayType>(po->allocatedType());
188 JLM_ASSERT(at);
189 // detect loads and stores attached to alloca
190 TraceAllocaUses ta(node->output(0));
191 // create memory + response
192 auto mem_outs = LocalMemoryOperation::create(at, node->region());
193 auto resp_outs = LocalMemoryResponseOperation::create(*mem_outs[0], ta.load_nodes.size());
194 std::cout << "alloca converted " << at->debug_string() << std::endl;
195 // replace gep outputs (convert pointer to index calculation)
196 // replace loads and stores
197 std::vector<jlm::rvsdg::Output *> load_addrs;
198 for (auto l : ta.load_nodes)
199 {
200 auto index = gep_to_index(l->input(0)->origin());
201 auto response = route_response_rhls(l->region(), resp_outs.front());
202 resp_outs.erase(resp_outs.begin());
203 std::vector<jlm::rvsdg::Output *> states;
204 for (size_t i = 1; i < l->ninputs(); ++i)
205 {
206 states.push_back(l->input(i)->origin());
207 }
208 auto load_outs = LocalLoadOperation::create(*index, states, *response);
209 auto nn = dynamic_cast<rvsdg::NodeOutput *>(load_outs[0])->node();
210 for (size_t i = 0; i < l->noutputs(); ++i)
211 {
212 l->output(i)->divert_users(nn->output(i));
213 }
214 remove(l);
215 auto addr = route_request_rhls(node->region(), load_outs.back());
216 load_addrs.push_back(addr);
217 }
218 std::vector<jlm::rvsdg::Output *> store_operands;
219 for (auto s : ta.store_nodes)
220 {
221 auto index = gep_to_index(s->input(0)->origin());
222 std::vector<jlm::rvsdg::Output *> states;
223 for (size_t i = 2; i < s->ninputs(); ++i)
224 {
225 states.push_back(s->input(i)->origin());
226 }
227 auto store_outs = LocalStoreOperation::create(*index, *s->input(1)->origin(), states);
228 auto nn = dynamic_cast<rvsdg::NodeOutput *>(store_outs[0])->node();
229 for (size_t i = 0; i < s->noutputs(); ++i)
230 {
231 s->output(i)->divert_users(nn->output(i));
232 }
233 remove(s);
234 auto addr = route_request_rhls(node->region(), store_outs[store_outs.size() - 2]);
235 auto data = route_request_rhls(node->region(), store_outs.back());
236 store_operands.push_back(addr);
237 store_operands.push_back(data);
238 }
239 // TODO: ensure that loads/stores are either alloca or global, never both
240 // TODO: ensure that loads/stores have same width and alignment and geps can be merged -
241 // otherwise slice? create request
242 auto req_outs = LocalMemoryRequestOperation::create(*mem_outs[1], load_addrs, store_operands);
243
244 // remove alloca from memstate merge
245 // TODO: handle general case of other nodes getting state edge without a merge
246 JLM_ASSERT(node->output(1)->nusers() == 1);
247 auto & merge_in = *node->output(1)->Users().begin();
248 auto merge_node = rvsdg::TryGetOwnerNode<rvsdg::Node>(merge_in);
249 if (dynamic_cast<const llvm::MemoryStateMergeOperation *>(&merge_node->GetOperation()))
250 {
251 // merge after alloca -> remove merge
252 JLM_ASSERT(merge_node->ninputs() == 2);
253 auto other_index = merge_in.index() ? 0 : 1;
254 merge_node->output(0)->divert_users(merge_node->input(other_index)->origin());
255 jlm::rvsdg::remove(merge_node);
256 }
257 else
258 {
259 // TODO: fix this properly by adding a state edge to the LambdaEntryMemState and routing it
260 // to the region
261 JLM_ASSERT(false);
262 }
263
264 // TODO: run dne to
265 // remove loads/stores
266 // remove geps
267 // remove alloca pointer users
268 // remove alloca
269 }
270 }
271}
272
273AllocaNodeConversion::~AllocaNodeConversion() noexcept = default;
274
278
279void
280AllocaNodeConversion::Run(rvsdg::RvsdgModule & rvsdgModule, util::StatisticsCollector &)
281{
282 alloca_conv(&rvsdgModule.Rvsdg().GetRootRegion());
283}
284
285} // namespace jlm::hls
void trace(jlm::rvsdg::Output *op)
std::vector< jlm::rvsdg::SimpleNode * > store_nodes
std::unordered_set< jlm::rvsdg::Output * > visited
TraceAllocaUses(jlm::rvsdg::Output *op)
std::vector< jlm::rvsdg::SimpleNode * > load_nodes
Conditional operator / pattern matching.
Definition gamma.hpp:99
Region & GetRootRegion() const noexcept
Definition graph.hpp:99
void divert_users(jlm::rvsdg::Output *new_origin)
Definition node.hpp:301
Represent acyclic RVSDG subgraphs.
Definition region.hpp:213
Graph & Rvsdg() noexcept
Represents an RVSDG transformation.
#define JLM_ASSERT(x)
Definition common.hpp:16
rvsdg::Output * route_response_rhls(rvsdg::Region *target, rvsdg::Output *response)
static jlm::rvsdg::Output * gep_to_index(jlm::rvsdg::Output *o)
rvsdg::Output * route_request_rhls(rvsdg::Region *target, rvsdg::Output *request)
static void alloca_conv(rvsdg::Region *region)
void MatchTypeOrFail(T &obj, const Fns &... fns)
Pattern match over subclass type of given object.
static void remove(Node *node)
Definition region.hpp:1035
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