6#include <gtest/gtest.h>
23#include <mlir/IR/OwningOpRef.h>
40 circt::firrtl::CircuitOp
55 std::unique_ptr<LlvmRvsdgModule>
Module_{};
58 template<
typename OpT>
66 if (::mlir::isa<OpT>(
op))
72 template<
typename OpT>
96 const std::vector<std::shared_ptr<const Type>> &
inputs,
97 const std::vector<std::shared_ptr<const Type>> &
outputs)
101 Module_->Rvsdg().GetRootRegion(),
106 template<
typename OpT>
139 return dynamic_cast<const MatchOperation *
>(&node.GetOperation());
151 auto *
matchOp = CreateMatchOp({ { 0, 0 }, { 1, 1 }, { 2, 2 } }, 3);
157 auto *
matchOp = CreateMatchOp({ { 0, 1 }, { 1, 0 } }, 2);
163 auto *
matchOp = CreateMatchOp({}, 2);
169 auto *
matchOp = CreateMatchOp({ { 0, 0 }, { 1, 2 } }, 3);
179 auto *
matchOp = CreateMatchOp({ { 0, 1 }, { 1, 3 }, { 2, 0 } }, 4);
197 auto *
matchOp = CreateMatchOp({ { 0, 1 } }, 4);
203 auto *
matchOp = CreateMatchOp({ { 0, 0 }, { 1, 1 } }, 16);
218 auto *
matchOp = CreateMatchOp({ { 0, 10 }, { 1, 20 }, { 2, 30 } }, 3);
219 std::unordered_map<uint64_t, uint64_t>
collected;
230 auto *
matchOp1 = CreateMatchOp({ { 0, 1 }, { 1, 0 } }, 3);
231 auto *
matchOp2 = CreateMatchOp({ { 0, 1 }, { 1, 0 } }, 3);
237 auto *
matchOp1 = CreateMatchOp({ { 0, 0 }, { 1, 1 } }, 2);
238 auto *
matchOp2 = CreateMatchOp({ { 0, 1 }, { 1, 0 } }, 2);
248 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
249 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
251 Lambda_->finalize({ addNode.output(0) });
258 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
259 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
261 Lambda_->finalize({
subNode.output(0) });
268 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
269 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
271 Lambda_->finalize({
mulNode.output(0) });
278 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
279 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
281 Lambda_->finalize({
andNode.output(0) });
288 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
289 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
291 Lambda_->finalize({
orNode.output(0) });
298 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
299 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
301 Lambda_->finalize({
xorNode.output(0) });
312 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
313 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
315 Lambda_->finalize({
shlNode.output(0) });
322 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
323 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
325 Lambda_->finalize({
lshrNode.output(0) });
336 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
337 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
339 Lambda_->finalize({
sdivNode.output(0) });
342 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
350 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
351 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
353 Lambda_->finalize({
ashrNode.output(0) });
356 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
364 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
365 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
367 Lambda_->finalize({
sremNode.output(0) });
370 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
384 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
385 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
387 Lambda_->finalize({
eqNode.output(0) });
396 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
397 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
399 Lambda_->finalize({
neqNode.output(0) });
408 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
409 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
411 Lambda_->finalize({
sgtNode.output(0) });
414 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
423 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
424 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
426 Lambda_->finalize({
sltNode.output(0) });
429 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
438 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
439 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
441 Lambda_->finalize({
sleNode.output(0) });
444 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
453 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
454 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
456 Lambda_->finalize({
sgeNode.output(0) });
459 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
468 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
469 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
471 Lambda_->finalize({
ultNode.output(0) });
480 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
481 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
483 Lambda_->finalize({
uleNode.output(0) });
492 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
493 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
495 Lambda_->finalize({
ugtNode.output(0) });
504 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
505 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
507 Lambda_->finalize({
ugeNode.output(0) });
520 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
531 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
536 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
564 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
576 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
581 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
589 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
594 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
607 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
612 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
629 Module_->Rvsdg().GetRootRegion(),
637 CreateMatchLambda(32, 4);
639 auto & predicate = *Lambda_->GetFunctionArguments()[0];
642 std::unordered_map<uint64_t, uint64_t>{ { 0, 3 }, { 1, 2 }, { 2, 1 }, { 3, 0 } },
645 Lambda_->finalize({ node.output(0) });
648 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
655 CreateMatchLambda(32, 4);
657 auto & predicate = *Lambda_->GetFunctionArguments()[0];
660 std::unordered_map<uint64_t, uint64_t>{ { 0, 5 }, { 2, 3 } },
663 Lambda_->finalize({ node.output(0) });
666 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
674 CreateMatchLambda(32, 4);
676 auto & predicate = *Lambda_->GetFunctionArguments()[0];
679 Lambda_->finalize({ node.output(0) });
682 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
690 Module_->Rvsdg().GetRootRegion(),
693 auto & predicate = *Lambda_->GetFunctionArguments()[0];
696 Lambda_->finalize({ node.output(0) });
699 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
719 Module_->Rvsdg().GetRootRegion(),
730 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
744 Module_->Rvsdg().GetRootRegion(),
747 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
748 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
754 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
764 Module_->Rvsdg().GetRootRegion(),
767 auto &
arg0 = *Lambda_->GetFunctionArguments()[0];
768 auto &
arg1 = *Lambda_->GetFunctionArguments()[1];
770 *Lambda_->subregion(),
776 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
789 Module_->Rvsdg().GetRootRegion(),
793 Lambda_->finalize({
testNode->output(0) });
802 catch (
const std::logic_error &)
820 Module_->Rvsdg().GetRootRegion(),
823 auto & predicate = *Lambda_->GetFunctionArguments()[0];
824 auto &
value0 = *Lambda_->GetFunctionArguments()[1];
825 auto &
value1 = *Lambda_->GetFunctionArguments()[2];
830 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
841 Module_->Rvsdg().GetRootRegion(),
844 auto & predicate = *Lambda_->GetFunctionArguments()[0];
845 auto &
value0 = *Lambda_->GetFunctionArguments()[1];
846 auto &
value1 = *Lambda_->GetFunctionArguments()[2];
847 auto &
value2 = *Lambda_->GetFunctionArguments()[3];
852 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
868 Module_->Rvsdg().GetRootRegion(),
871 auto & ptr = *Lambda_->GetFunctionArguments()[0];
872 auto & index = *Lambda_->GetFunctionArguments()[1];
877 mlir::OwningOpRef<circt::firrtl::CircuitOp>
circuit(
converter.TestMlirGen(Lambda_));
TEST_F(IdentityMappingTest, IdentityMapping)
LambdaNode * CreateLambda(const std::vector< std::shared_ptr< const Type > > &inputs, const std::vector< std::shared_ptr< const Type > > &outputs)
LambdaNode * CreateMatchLambda(int predicateBits, int outBits)
std::unique_ptr< LlvmRvsdgModule > Module_
bool AssertFirrtlOpExists(mlir::OwningOpRef< circt::firrtl::CircuitOp > &circuit)
bool AssertFirrtlOpExists(mlir::Operation *circuit)
const MatchOperation * CreateMatchOp(const std::unordered_map< uint64_t, uint64_t > &mapping, uint64_t nalternatives)
std::unique_ptr< LlvmRvsdgModule > Module_
TestableRhlsToFirrtlConverter Converter_
static std::vector< jlm::rvsdg::Output * > create(jlm::rvsdg::Output &predicate, const std::vector< jlm::rvsdg::Output * > &alternatives, bool discarding, bool loop=false)
bool IsIdentityMapping(const rvsdg::MatchOperation &op)
circt::firrtl::CircuitOp MlirGen(const rvsdg::LambdaNode *lamdaNode)
bool TestIsIdentityMapping(const rvsdg::MatchOperation &op)
circt::firrtl::CircuitOp TestMlirGen(const rvsdg::LambdaNode *lambdaNode)
static std::shared_ptr< const ArrayType > Create(std::shared_ptr< const Type > type, size_t nelements)
static jlm::rvsdg::Output * create(jlm::rvsdg::Output *operand, std::shared_ptr< const jlm::rvsdg::Type > rtype)
static rvsdg::Output * create(rvsdg::Output *baseAddress, const std::vector< rvsdg::Output * > &indices, std::shared_ptr< const rvsdg::Type > gepType)
static std::unique_ptr< ThreeAddressCode > create(const Variable *argument)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & Create(rvsdg::Region ®ion, IntegerValueRepresentation representation)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::Node & createNode(const size_t numBits, rvsdg::Output &operand1, rvsdg::Output &operand2)
static rvsdg::SimpleNode & CreateNode(rvsdg::Region ®ion, const std::vector< rvsdg::Output * > &operands, const std::vector< MemoryNodeId > &memoryNodeIds)
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)
static std::unique_ptr< LlvmRvsdgModule > Create(const util::FilePath &sourceFileName, const std::string &targetTriple, const std::string &dataLayout)
static rvsdg::Output * Create(const std::vector< rvsdg::Output * > &operands)
static std::shared_ptr< const MemoryStateType > Create()
static std::shared_ptr< const PointerType > Create()
static rvsdg::Output & create(size_t ndstbits, rvsdg::Output &operand)
static rvsdg::Output & create(size_t ndstbits, rvsdg::Output &operand)
static jlm::rvsdg::Output * Create(rvsdg::Region ®ion, std::shared_ptr< const jlm::rvsdg::Type > type)
static rvsdg::Output & create(size_t ndstbits, rvsdg::Output &operand)
static std::shared_ptr< const BitType > Create(std::size_t nbits)
Creates bit type of specified width.
static Output & create(Region ®ion, ControlValueRepresentation value)
static std::shared_ptr< const ControlType > Create(std::size_t nalternatives)
Instantiates control type.
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 LambdaNode * Create(rvsdg::Region &parent, std::unique_ptr< LambdaOperation > operation)
static Node & CreateNode(Output &predicate, const std::unordered_map< uint64_t, uint64_t > &mapping, const uint64_t defaultAlternative, const size_t numAlternatives)
NodeOutput * output(size_t index) const noexcept
static SimpleNode * createNode(Region *region, const std::vector< Output * > &operands, std::vector< std::shared_ptr< const Type > > resultTypes)
Global memory state passed between functions.
static std::vector< jlm::rvsdg::Output * > outputs(const Node *node)
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.