diff options
Diffstat (limited to 'src/armnn/test/TestNameAndDescriptorLayerVisitor.hpp')
-rw-r--r-- | src/armnn/test/TestNameAndDescriptorLayerVisitor.hpp | 34 |
1 files changed, 34 insertions, 0 deletions
diff --git a/src/armnn/test/TestNameAndDescriptorLayerVisitor.hpp b/src/armnn/test/TestNameAndDescriptorLayerVisitor.hpp index f1936d6847..aa0b3597fa 100644 --- a/src/armnn/test/TestNameAndDescriptorLayerVisitor.hpp +++ b/src/armnn/test/TestNameAndDescriptorLayerVisitor.hpp @@ -418,6 +418,40 @@ public: }; }; +class TestInstanceNormalizationLayerVisitor : public TestLayerVisitor +{ +private: + InstanceNormalizationDescriptor m_VisitorDescriptor; + +public: + explicit TestInstanceNormalizationLayerVisitor(const InstanceNormalizationDescriptor& desc, + const char* name = nullptr) + : TestLayerVisitor(name) + { + m_VisitorDescriptor.m_Beta = desc.m_Beta; + m_VisitorDescriptor.m_Gamma = desc.m_Gamma; + m_VisitorDescriptor.m_Eps = desc.m_Eps; + m_VisitorDescriptor.m_DataLayout = desc.m_DataLayout; + }; + + void CheckDescriptor(const InstanceNormalizationDescriptor& desc) + { + BOOST_CHECK(desc.m_Beta == m_VisitorDescriptor.m_Beta); + BOOST_CHECK(desc.m_Gamma == m_VisitorDescriptor.m_Gamma); + BOOST_CHECK(desc.m_Eps == m_VisitorDescriptor.m_Eps); + BOOST_CHECK(desc.m_DataLayout == m_VisitorDescriptor.m_DataLayout); + } + + void VisitInstanceNormalizationLayer(const IConnectableLayer* layer, + const InstanceNormalizationDescriptor& desc, + const char* name = nullptr) override + { + CheckLayerPointer(layer); + CheckDescriptor(desc); + CheckLayerName(name); + }; +}; + class TestL2NormalizationLayerVisitor : public TestLayerVisitor { private: |