Jlm
Loading...
Searching...
No Matches
lambda.cpp
Go to the documentation of this file.
1/*
2 * Copyright 2018 Nico Reißmann <nico.reissmann@gmail.com>
3 * Copyright 2025 Helge Bahmann <hcb@chaoticmind.net>
4 * See COPYING for terms of redistribution.
5 */
6
9#include <jlm/util/strfmt.hpp>
10
11namespace jlm::rvsdg
12{
13
15
16LambdaOperation::LambdaOperation(std::shared_ptr<const FunctionType> type)
17 : type_(std::move(type))
18{}
19
20std::string
22{
23 return util::strfmt("Lambda[", Type()->debug_string(), "]");
24}
25
26bool
28{
29 auto op = dynamic_cast<const LambdaOperation *>(&other);
30 return op && op->type() == type();
31}
32
33std::unique_ptr<rvsdg::Operation>
35{
36 return std::make_unique<LambdaOperation>(*this);
37}
38
39LambdaNode::~LambdaNode() = default;
40
41LambdaNode::LambdaNode(rvsdg::Region & parent, std::unique_ptr<LambdaOperation> op)
42 : StructuralNode(&parent, 1),
43 Operation_(std::move(op))
44{
45 for (auto & argumentType : GetOperation().Type()->Arguments())
46 {
48 }
49}
50
56
57[[nodiscard]] std::vector<rvsdg::Output *>
59{
60 std::vector<rvsdg::Output *> arguments;
61 const auto & type = GetOperation().Type();
62 for (std::size_t n = 0; n < type->Arguments().size(); ++n)
63 {
64 arguments.push_back(subregion()->argument(n));
65 }
66 return arguments;
67}
68
69[[nodiscard]] std::vector<rvsdg::Input *>
71{
72 std::vector<rvsdg::Input *> results;
73 for (std::size_t n = 0; n < subregion()->nresults(); ++n)
74 {
75 results.push_back(subregion()->result(n));
76 }
77 return results;
78}
79
81LambdaNode::MapInputContextVar(const rvsdg::Input & input) const noexcept
82{
84 return ContextVar{ const_cast<rvsdg::Input *>(&input),
85 subregion()->argument(GetOperation().Type()->NumArguments() + input.index()) };
86}
87
88[[nodiscard]] std::optional<LambdaNode::ContextVar>
89LambdaNode::MapBinderContextVar(const rvsdg::Output & output) const noexcept
90{
91 JLM_ASSERT(rvsdg::TryGetOwnerRegion(output) == subregion());
92 auto numArguments = GetOperation().Type()->NumArguments();
93 if (output.index() >= numArguments)
94 {
95 return ContextVar{ input(output.index() - GetOperation().Type()->NumArguments()),
96 const_cast<rvsdg::Output *>(&output) };
97 }
98 else
99 {
100 return std::nullopt;
101 }
102}
103
104std::variant<LambdaNode::ArgumentVar, LambdaNode::ContextVar>
106{
108 std::size_t nargs = GetOperation().Type()->NumArguments();
109 if (output.index() < nargs)
110 {
111 return ArgumentVar{ subregion()->argument(output.index()) };
112 }
113 else
114 {
116 }
117}
118
119[[nodiscard]] std::vector<LambdaNode::ContextVar>
121{
122 std::vector<ContextVar> vars;
123 for (size_t n = 0; n < ninputs(); ++n)
124 {
125 vars.push_back(
126 ContextVar{ input(n), subregion()->argument(n + GetOperation().Type()->NumArguments()) });
127 }
128 return vars;
129}
130
133{
134 const auto input =
135 addInput(std::make_unique<StructuralInput>(this, &origin, origin.Type()), true);
136 const auto argument = &RegionArgument::Create(*subregion(), input, origin.Type());
137 return ContextVar{ input, argument };
138}
139
141LambdaNode::Create(rvsdg::Region & parent, std::unique_ptr<LambdaOperation> operation)
142{
143 return new LambdaNode(parent, std::move(operation));
144}
145
147LambdaNode::finalize(const std::vector<jlm::rvsdg::Output *> & results)
148{
149 /* check if finalized was already called */
150 if (noutputs() > 0)
151 {
152 JLM_ASSERT(noutputs() == 1);
153 return output();
154 }
155
156 if (GetOperation().type().NumResults() != results.size())
157 throw util::Error("Incorrect number of results.");
158
159 for (size_t n = 0; n < results.size(); n++)
160 {
161 auto & expected = GetOperation().type().ResultType(n);
162 auto & received = *results[n]->Type();
163 if (*results[n]->Type() != GetOperation().type().ResultType(n))
164 throw util::Error("Expected " + expected.debug_string() + ", got " + received.debug_string());
165
166 if (results[n]->region() != subregion())
167 throw util::Error("Invalid operand region.");
168 }
169
170 for (const auto & origin : results)
171 rvsdg::RegionResult::Create(*origin->region(), *origin, nullptr, origin->Type());
172
173 return addOutput(std::make_unique<StructuralOutput>(this, GetOperation().Type()));
174}
175
181
183LambdaNode::copy(rvsdg::Region * region, const std::vector<jlm::rvsdg::Output *> & operands) const
184{
185 return util::assertedCast<LambdaNode>(rvsdg::Node::copy(region, operands));
186}
187
190{
191 const auto & op = GetOperation();
192 auto lambda = Create(
193 *region,
194 std::unique_ptr<LambdaOperation>(util::assertedCast<LambdaOperation>(op.copy().release())));
195
196 /* add context variables */
198 for (const auto & cv : GetContextVars())
199 {
200 auto origin = &smap.lookup(*cv.input->origin());
201 subregionmap.insert(cv.inner, lambda->AddContextVar(*origin).inner);
202 }
203
204 /* collect function arguments */
205 auto args = GetFunctionArguments();
206 auto newArgs = lambda->GetFunctionArguments();
207 JLM_ASSERT(args.size() == newArgs.size());
208 for (std::size_t n = 0; n < args.size(); ++n)
209 {
210 subregionmap.insert(args[n], newArgs[n]);
211 }
212
213 /* copy subregion */
214 subregion()->copy(lambda->subregion(), subregionmap);
215
216 /* collect function results */
217 std::vector<jlm::rvsdg::Output *> results;
218 for (auto result : GetFunctionResults())
219 results.push_back(&subregionmap.lookup(*result->origin()));
220
221 /* finalize lambda */
222 auto o = lambda->finalize(results);
223 smap.insert(output(), o);
224
225 return lambda;
226}
227
228LambdaBuilder::LambdaBuilder(Region & region, std::vector<std::shared_ptr<const Type>> argtypes)
229 : Node_(LambdaNode::Create(
230 region,
231 std::make_unique<LambdaOperation>(FunctionType::Create(std::move(argtypes), {}))))
232{
233 // Note that the above inserts a "placeholder" function type, for now.
234 // This is to avoid requiring the caller to specify the return type(s)
235 // already when starting to construct the object. It is sometimes easier
236 // to let them be determined while building.
237}
238
239std::vector<Output *>
245
252
259
260Output &
262 const std::vector<jlm::rvsdg::Output *> & results,
263 std::unique_ptr<LambdaOperation> op)
264{
266 Node_->Operation_ = std::move(op);
267 auto output = Node_->finalize(results);
268 Node_ = nullptr;
269 return *output;
270}
271
274{
275 auto it = &node;
276 while (it)
277 {
278 if (auto lambda = dynamic_cast<rvsdg::LambdaNode *>(it))
279 return *lambda;
280 it = it->region()->node();
281 }
282 throw std::logic_error("node was not in a lambda");
283}
284
287{
288 return getSurroundingLambdaNode(const_cast<rvsdg::Node &>(node));
289}
290
291}
std::int64_t expected
util::HashSet< rvsdg::Output * > arguments
Function type class.
const jlm::rvsdg::Type & ResultType(size_t index) const noexcept
Output & Finalize(const std::vector< jlm::rvsdg::Output * > &results, std::unique_ptr< LambdaOperation > op)
Verifies well-formedness of lambda node and completes it.
Definition lambda.cpp:261
LambdaNode::ContextVar AddContextVar(jlm::rvsdg::Output &origin)
Adds a context/free variable to the lambda node.
Definition lambda.cpp:254
rvsdg::Region * GetRegion() noexcept
Returns region to place nodes in.
Definition lambda.cpp:247
LambdaBuilder(Region &region, std::vector< std::shared_ptr< const Type > > argtypes)
Creates builder for a lambda construct.
Definition lambda.cpp:228
std::vector< Output * > Arguments()
Obtains definition points of parameters to the function.
Definition lambda.cpp:240
LambdaNode * copy(rvsdg::Region *region, const std::vector< jlm::rvsdg::Output * > &operands) const override
Definition lambda.cpp:183
rvsdg::Output * finalize(const std::vector< jlm::rvsdg::Output * > &results)
Definition lambda.cpp:147
std::variant< ArgumentVar, ContextVar > MapArgument(const rvsdg::Output &output) const
Maps region argument to its disposition (formal argument or context var).
Definition lambda.cpp:105
std::vector< rvsdg::Output * > GetFunctionArguments() const
Definition lambda.cpp:58
ContextVar MapInputContextVar(const rvsdg::Input &input) const noexcept
Maps input to context variable.
Definition lambda.cpp:81
std::optional< ContextVar > MapBinderContextVar(const rvsdg::Output &output) const noexcept
Maps bound variable reference to context variable.
Definition lambda.cpp:89
ContextVar AddContextVar(jlm::rvsdg::Output &origin)
Adds a context/free variable to the lambda node.
Definition lambda.cpp:132
LambdaNode(rvsdg::Region &parent, std::unique_ptr< LambdaOperation > op)
Definition lambda.cpp:41
std::unique_ptr< LambdaOperation > Operation_
Definition lambda.hpp:289
static LambdaNode * Create(rvsdg::Region &parent, std::unique_ptr< LambdaOperation > operation)
Definition lambda.cpp:141
std::vector< rvsdg::Input * > GetFunctionResults() const
Definition lambda.cpp:70
rvsdg::Output * output() const noexcept
Definition lambda.cpp:177
rvsdg::Region * subregion() const noexcept
Definition lambda.hpp:138
std::vector< ContextVar > GetContextVars() const noexcept
Gets all bound context variables.
Definition lambda.cpp:120
LambdaOperation & GetOperation() const noexcept override
Definition lambda.cpp:52
Lambda operation.
Definition lambda.hpp:29
LambdaOperation(std::shared_ptr< const FunctionType > type)
Definition lambda.cpp:16
bool operator==(const Operation &other) const noexcept override
Definition lambda.cpp:27
const FunctionType & type() const noexcept
Definition lambda.hpp:36
const std::shared_ptr< const FunctionType > & Type() const noexcept
Definition lambda.hpp:42
std::unique_ptr< Operation > copy() const override
Definition lambda.cpp:34
std::string debug_string() const override
Definition lambda.cpp:21
rvsdg::Region * region() const noexcept
Definition node.hpp:761
size_t ninputs() const noexcept
Definition node.hpp:609
size_t noutputs() const noexcept
Definition node.hpp:644
virtual Node * copy(rvsdg::Region *region, const std::vector< jlm::rvsdg::Output * > &operands) const
Definition node.cpp:369
size_t index() const noexcept
Definition node.hpp:274
const std::shared_ptr< const rvsdg::Type > & Type() const noexcept
Definition node.hpp:366
static RegionArgument & Create(rvsdg::Region &region, StructuralInput *input, std::shared_ptr< const rvsdg::Type > type)
Creates region entry argument.
Definition region.cpp:63
static RegionResult & Create(rvsdg::Region &region, rvsdg::Output &origin, StructuralOutput *output, std::shared_ptr< const rvsdg::Type > type)
Create region exit result.
Definition region.cpp:112
Represent acyclic RVSDG subgraphs.
Definition region.hpp:213
RegionArgument * argument(size_t index) const noexcept
Definition region.hpp:466
void copy(Region *target, SubstitutionMap &smap) const
Copy a region with substitutions.
Definition region.cpp:317
size_t nresults() const noexcept
Definition region.hpp:494
RegionResult * result(size_t index) const noexcept
Definition region.hpp:500
rvsdg::StructuralNode * node() const noexcept
Definition region.hpp:301
StructuralInput * addInput(std::unique_ptr< StructuralInput > input, bool notifyRegion)
StructuralOutput * addOutput(std::unique_ptr< StructuralOutput > input)
StructuralOutput * output(size_t index) const noexcept
StructuralInput * input(size_t index) const noexcept
Output & lookup(const Output &original) const
constexpr Type() noexcept
Definition type.hpp:46
#define JLM_ASSERT(x)
Definition common.hpp:16
static std::vector< jlm::rvsdg::Output * > operands(const Node *node)
Definition node.hpp:1049
Region * TryGetOwnerRegion(const rvsdg::Input &input) noexcept
Definition node.hpp:1021
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.
Definition node.hpp:872
rvsdg::LambdaNode & getSurroundingLambdaNode(rvsdg::Node &node)
Definition lambda.cpp:273
static std::string strfmt(Args... args)
Definition strfmt.hpp:35
Formal argument variable.
Definition lambda.hpp:124
Bound context variable.
Definition lambda.hpp:100