Comparison of InputDescription and OutputDescription in SubGraph (#4123)
* draft - looks good, need merge with other branch * apply suggestions from PR * apply suggestions from PR * add unit test - NotValidInputOrOutput * rephraze equal for constant * cleanup all TODOs * Try to fix TI * try to satisfy PR comments * fix spelling Co-authored-by: Patryk Elszkowski <patryk.elszkowki@intel.com>
This commit is contained in:
parent
8f4d2be537
commit
af712743ab
|
|
@ -43,11 +43,7 @@ TEST_P(SerializationTest, CompareFunctions) {
|
|||
InferenceEngine::Core ie;
|
||||
InferenceEngine::CNNNetwork expected;
|
||||
|
||||
if (!m_binary_path.empty()) {
|
||||
expected = ie.ReadNetwork(m_model_path, m_binary_path);
|
||||
} else {
|
||||
expected = ie.ReadNetwork(m_model_path);
|
||||
}
|
||||
expected = ie.ReadNetwork(m_model_path, m_binary_path);
|
||||
expected.serialize(m_out_xml_path, m_out_bin_path);
|
||||
auto result = ie.ReadNetwork(m_out_xml_path, m_out_bin_path);
|
||||
|
||||
|
|
|
|||
|
|
@ -6,16 +6,17 @@
|
|||
|
||||
#include <gmock/gmock-matchers.h>
|
||||
|
||||
#include "common_test_utils/test_common.hpp"
|
||||
#include <string>
|
||||
#include <memory>
|
||||
#include <queue>
|
||||
#include <string>
|
||||
#include "common_test_utils/test_common.hpp"
|
||||
|
||||
#include <ngraph/pass/manager.hpp>
|
||||
#include <ngraph/function.hpp>
|
||||
#include <ngraph/opsets/opset5.hpp>
|
||||
#include <transformations/utils/utils.hpp>
|
||||
#include <ngraph/opsets/opset6.hpp>
|
||||
#include <ngraph/pass/manager.hpp>
|
||||
#include <transformations/init_node_info.hpp>
|
||||
#include <transformations/utils/utils.hpp>
|
||||
|
||||
#include "common_test_utils/ngraph_test_utils.hpp"
|
||||
|
||||
|
|
@ -207,9 +208,141 @@ TEST(TransformationTests, CompareFunctoinsTINegative) {
|
|||
ngraph::ParameterVector{X, Y, Z});
|
||||
}
|
||||
|
||||
auto res = compare_functions(f, f_ref);
|
||||
EXPECT_FALSE(res.first);
|
||||
EXPECT_EQ(res.second, "LSTMCell/4 != Relu/0");
|
||||
const auto fc = FunctionsComparator::with_default().enable(FunctionsComparator::ATTRIBUTES);
|
||||
auto res = fc(f, f_ref);
|
||||
EXPECT_FALSE(res.valid);
|
||||
EXPECT_THAT(res.message, HasSubstr("LSTMCell/4 != Relu/0"));
|
||||
}
|
||||
|
||||
TEST(TransformationTests, CompareFunctoinsTINegativeDifferentElementTypeBetweenSubGraphsInputs) {
|
||||
const auto createFunc = [](element::Type e_type) {
|
||||
using namespace opset6;
|
||||
|
||||
auto X = std::make_shared<Parameter>(e_type, Shape{1, 2});
|
||||
auto Y = std::make_shared<Parameter>(e_type, Shape{1, 2});
|
||||
|
||||
auto Xi = std::make_shared<Parameter>(e_type, Shape{1, 2});
|
||||
auto Yi = std::make_shared<Parameter>(e_type, Shape{1, 2});
|
||||
|
||||
// Body
|
||||
auto add = std::make_shared<Add>(Xi, Yi);
|
||||
auto result = std::make_shared<Result>(add);
|
||||
|
||||
auto ti_body = std::make_shared<Function>(OutputVector{result}, ParameterVector{Xi, Yi});
|
||||
|
||||
auto ti = std::make_shared<TensorIterator>();
|
||||
ti->set_body(ti_body);
|
||||
ti->set_sliced_input(Xi, X, 0, 1, 1, -1, 1);
|
||||
ti->set_sliced_input(Yi, Y, 0, 1, 1, -1, 1);
|
||||
|
||||
auto out = ti->get_concatenated_slices(result, 0, 1, 1, -1, 1);
|
||||
|
||||
return std::make_shared<Function>(
|
||||
NodeVector{out.get_node_shared_ptr()}, ParameterVector{X, Y});
|
||||
};
|
||||
const auto f1 = createFunc(element::f32);
|
||||
const auto f2 = createFunc(element::f16);
|
||||
|
||||
using FnCmp = FunctionsComparator;
|
||||
|
||||
const auto result = FnCmp::with_default().compare(f1, f2);
|
||||
|
||||
EXPECT_FALSE(result.valid);
|
||||
EXPECT_THAT(result.message, HasSubstr("different SubGraph InputDescription"));
|
||||
}
|
||||
|
||||
TEST(TransformationTests, CompareFunctoinsTINegativeDifferentElementTypeBetweenInputAndParameter) {
|
||||
const auto createFunc = [](element::Type e_type) {
|
||||
using namespace opset6;
|
||||
|
||||
auto X = std::make_shared<Parameter>(element::f64, Shape{1, 2}); // <<
|
||||
auto Y = std::make_shared<Parameter>(e_type, Shape{1, 2});
|
||||
|
||||
auto Xi = std::make_shared<Parameter>(e_type, Shape{1, 2});
|
||||
auto Yi = std::make_shared<Parameter>(e_type, Shape{1, 2});
|
||||
|
||||
// Body
|
||||
auto add = std::make_shared<Add>(Xi, Yi);
|
||||
auto result = std::make_shared<Result>(add);
|
||||
|
||||
auto ti_body = std::make_shared<Function>(OutputVector{result}, ParameterVector{Xi, Yi});
|
||||
|
||||
auto ti = std::make_shared<TensorIterator>();
|
||||
ti->set_body(ti_body);
|
||||
ti->set_sliced_input(Xi, X, 0, 1, 1, -1, 1);
|
||||
ti->set_sliced_input(Yi, Y, 0, 1, 1, -1, 1);
|
||||
|
||||
auto out = ti->get_concatenated_slices(result, 0, 1, 1, -1, 1);
|
||||
|
||||
return std::make_shared<Function>(
|
||||
NodeVector{out.get_node_shared_ptr()}, ParameterVector{X, Y});
|
||||
};
|
||||
const auto f1 = createFunc(element::f32);
|
||||
|
||||
using FnCmp = FunctionsComparator;
|
||||
|
||||
const auto result = FnCmp::with_default().compare(f1, f1);
|
||||
|
||||
EXPECT_FALSE(result.valid);
|
||||
EXPECT_THAT(result.message, HasSubstr("inputs and parameters mismatch"));
|
||||
}
|
||||
|
||||
TEST(TransformationTests, CompareFunctoinsTINegativeDifferentElementTypeBetweentResultAndOutput) {
|
||||
const auto createFunc = [](element::Type result_element_type, const Shape& result_shape) {
|
||||
using namespace opset6;
|
||||
|
||||
auto X = std::make_shared<Parameter>(element::f32, Shape{1, 2});
|
||||
auto Y = std::make_shared<Parameter>(element::f32, Shape{1, 2});
|
||||
|
||||
auto Xi = std::make_shared<Parameter>(element::f32, Shape{1, 2});
|
||||
auto Yi = std::make_shared<Parameter>(element::f32, Shape{1, 2});
|
||||
|
||||
// Body
|
||||
auto add = std::make_shared<Add>(Xi, Yi);
|
||||
auto result = std::make_shared<Result>(add);
|
||||
|
||||
auto ti_body = std::make_shared<Function>(OutputVector{result}, ParameterVector{Xi, Yi});
|
||||
|
||||
auto ti = std::make_shared<TensorIterator>();
|
||||
ti->set_body(ti_body);
|
||||
ti->set_sliced_input(Xi, X, 0, 1, 1, -1, 1);
|
||||
ti->set_sliced_input(Yi, Y, 0, 1, 1, -1, 1);
|
||||
|
||||
auto out = ti->get_concatenated_slices(result, 0, 1, 1, -1, 1);
|
||||
|
||||
auto fn = std::make_shared<Function>(
|
||||
NodeVector{out.get_node_shared_ptr()}, ParameterVector{X, Y});
|
||||
|
||||
/// <<
|
||||
auto&& result_out = result->output(0);
|
||||
Node* result_out_node = result_out.get_node();
|
||||
result_out_node->set_output_type(0, result_element_type, result_shape);
|
||||
/// <<
|
||||
|
||||
return fn;
|
||||
};
|
||||
{ // check element type difference
|
||||
const auto f1 = createFunc(element::u16, Shape{10, 20});
|
||||
const auto f2 = createFunc(element::u64, Shape{10, 20});
|
||||
|
||||
using FnCmp = FunctionsComparator;
|
||||
|
||||
const auto result = FnCmp::with_default().compare(f1, f2);
|
||||
|
||||
EXPECT_FALSE(result.valid);
|
||||
EXPECT_THAT(result.message, HasSubstr("outputs and results mismatch"));
|
||||
}
|
||||
{ // check Shape difference
|
||||
const auto f1 = createFunc(element::u16, Shape{11, 20});
|
||||
const auto f2 = createFunc(element::u16, Shape{12, 20});
|
||||
|
||||
using FnCmp = FunctionsComparator;
|
||||
|
||||
const auto result = FnCmp::with_default().compare(f1, f2);
|
||||
|
||||
EXPECT_FALSE(result.valid);
|
||||
EXPECT_THAT(result.message, HasSubstr("outputs and results mismatch"));
|
||||
}
|
||||
}
|
||||
|
||||
TEST(TransformationTests, ConstantNegativeDifferentElementType) {
|
||||
|
|
|
|||
|
|
@ -4,7 +4,9 @@
|
|||
|
||||
#include "ngraph_test_utils.hpp"
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstdint>
|
||||
#include <iostream>
|
||||
#include <map>
|
||||
#include <memory>
|
||||
|
|
@ -22,7 +24,10 @@
|
|||
#include <ngraph/opsets/opset6.hpp>
|
||||
#include <ngraph/pass/visualize_tree.hpp>
|
||||
|
||||
#include "details/ie_exception.hpp"
|
||||
|
||||
namespace {
|
||||
inline namespace tools {
|
||||
bool isTypeRelaxed(const std::string& type) {
|
||||
return type.find_first_of("TypeRelaxed") == 0;
|
||||
}
|
||||
|
|
@ -74,7 +79,8 @@ bool less_by_name(
|
|||
|
||||
template <typename T>
|
||||
std::string to_str(const T& v) {
|
||||
return std::to_string(v);
|
||||
using std::to_string;
|
||||
return to_string(v);
|
||||
}
|
||||
|
||||
std::string typeInfoToStr(const ngraph::Node::type_info_t& typeInfo) {
|
||||
|
|
@ -95,6 +101,7 @@ std::string tensor_names(const ngraph::descriptor::Tensor& t) {
|
|||
}
|
||||
return "\"" + n + "\"";
|
||||
}
|
||||
} // namespace tools
|
||||
|
||||
class Comparator {
|
||||
public:
|
||||
|
|
@ -141,7 +148,9 @@ private:
|
|||
std::unordered_set<ngraph::Node*> used;
|
||||
};
|
||||
|
||||
namespace attr_comparison {
|
||||
namespace attributes {
|
||||
|
||||
namespace detail {
|
||||
|
||||
using AttrName = std::string;
|
||||
|
||||
|
|
@ -242,9 +251,9 @@ class Storage : private AttributeStorage<MemoryChunk>,
|
|||
private AttributeStorage<std::vector<float>>,
|
||||
private AttributeStorage<std::vector<double>>,
|
||||
private AttributeStorage<std::vector<std::string>>,
|
||||
private AttributeStorage<std::shared_ptr<ngraph::Function>>,
|
||||
private AttributeStorage<SubGraphOpInputDescription>,
|
||||
private AttributeStorage<SubGraphOpOutputDescription>,
|
||||
private AttributeStorage<SpecialBodyPorts> {
|
||||
private AttributeStorage<SubGraphOpOutputDescription> {
|
||||
public:
|
||||
template <typename AttrValue>
|
||||
const AttributeStorage<AttrValue>& storage() const {
|
||||
|
|
@ -280,9 +289,9 @@ public:
|
|||
storage<std::vector<float>>().get_attributes_number() +
|
||||
storage<std::vector<double>>().get_attributes_number() +
|
||||
storage<std::vector<std::string>>().get_attributes_number() +
|
||||
storage<std::shared_ptr<ngraph::Function>>().get_attributes_number() +
|
||||
storage<SubGraphOpInputDescription>().get_attributes_number() +
|
||||
storage<SubGraphOpOutputDescription>().get_attributes_number() +
|
||||
storage<SpecialBodyPorts>().get_attributes_number();
|
||||
storage<SubGraphOpOutputDescription>().get_attributes_number();
|
||||
}
|
||||
};
|
||||
|
||||
|
|
@ -298,9 +307,9 @@ public:
|
|||
auto outputs =
|
||||
ngraph::as_type<ngraph::AttributeAdapter<SubGraphOpOutputDescription>>(&adapter)) {
|
||||
insert(name, outputs->get());
|
||||
} else if (
|
||||
auto ports = ngraph::as_type<ngraph::AttributeAdapter<SpecialBodyPorts>>(&adapter)) {
|
||||
insert(name, ports->get());
|
||||
} else if (ngraph::is_type<ngraph::AttributeAdapter<SpecialBodyPorts>>(&adapter)) {
|
||||
// drop comparison, no more info than port indexes which will be check in
|
||||
// subgraph::compare_io
|
||||
} else if (
|
||||
auto a = ngraph::as_type<
|
||||
ngraph::AttributeAdapter<std::shared_ptr<ngraph::runtime::AlignedBuffer>>>(
|
||||
|
|
@ -343,14 +352,10 @@ public:
|
|||
ON_ADAPTER(std::vector<float>)
|
||||
ON_ADAPTER(std::vector<double>)
|
||||
ON_ADAPTER(std::vector<std::string>)
|
||||
ON_ADAPTER(std::shared_ptr<ngraph::Function>)
|
||||
|
||||
#undef ON_ADAPTER
|
||||
|
||||
void on_adapter(
|
||||
const std::string&, ngraph::ValueAccessor<std::shared_ptr<ngraph::Function>>&) override {
|
||||
// handled by `compare_functions` drop it here
|
||||
}
|
||||
|
||||
template <typename AttrValue>
|
||||
const AttrValue* get(const AttrName& name) const {
|
||||
return storage<AttrValue>().get_value(name);
|
||||
|
|
@ -510,13 +515,6 @@ struct Equal<SubGraphOpOutputDescription> {
|
|||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Equal<SpecialBodyPorts> {
|
||||
static bool equal_value(const SpecialBodyPorts& lhs, const SpecialBodyPorts& rhs) {
|
||||
return lhs.current_iteration_input_idx == rhs.current_iteration_input_idx;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Equal<uint8_t*> {
|
||||
static constexpr uint8_t BITS_IN_BYTE_COUNT = 8;
|
||||
|
|
@ -644,45 +642,11 @@ struct Get<
|
|||
|
||||
class ReadAndCompareAttributes : public ngraph::AttributeVisitor {
|
||||
public:
|
||||
ReadAndCompareAttributes(const ReadAndStoreAttributes& ref)
|
||||
: m_attr_ref(ref), m_cmp_result{ref.read_result()} {}
|
||||
ReadAndCompareAttributes(const ReadAndStoreAttributes& ref, Comparator::CmpValues check_flags)
|
||||
: m_attr_ref(ref), m_cmp_result{ref.read_result()}, m_check_flags(check_flags) {}
|
||||
|
||||
void on_adapter(const std::string& name, ngraph::ValueAccessor<void>& adapter) override {
|
||||
if (should_return()) {
|
||||
return;
|
||||
}
|
||||
m_visited_attributes.insert(name);
|
||||
if (auto inputs =
|
||||
ngraph::as_type<ngraph::AttributeAdapter<SubGraphOpInputDescription>>(&adapter)) {
|
||||
verify(name, inputs->get());
|
||||
} else if (
|
||||
auto outputs =
|
||||
ngraph::as_type<ngraph::AttributeAdapter<SubGraphOpOutputDescription>>(&adapter)) {
|
||||
verify(name, outputs->get());
|
||||
} else if (
|
||||
auto ports = ngraph::as_type<ngraph::AttributeAdapter<SpecialBodyPorts>>(&adapter)) {
|
||||
verify(name, ports->get());
|
||||
} else if (
|
||||
auto a = ngraph::as_type<
|
||||
ngraph::AttributeAdapter<std::shared_ptr<ngraph::runtime::AlignedBuffer>>>(
|
||||
&adapter)) {
|
||||
m_visited_attributes.insert(name);
|
||||
const auto ref_value = m_attr_ref.get<storage::MemoryChunk>(name);
|
||||
if (!ref_value) {
|
||||
m_cmp_result += "missing attribute name: '" + name + "'";
|
||||
return;
|
||||
}
|
||||
|
||||
if (a->get()->size() != ref_value->size() ||
|
||||
std::memcmp(ref_value->data(), a->get()->get_ptr(), ref_value->size()) != 0) {
|
||||
m_cmp_result += "mismatch in value: '" + name + "' : look in to the mem buffer";
|
||||
return;
|
||||
}
|
||||
} else {
|
||||
m_cmp_result += "compare attr [ ERR ]: " + name +
|
||||
" [drop `void` comparison which is '" + adapter.get_type_info().name +
|
||||
"']";
|
||||
}
|
||||
verify_others(name, adapter);
|
||||
}
|
||||
|
||||
#define ON_ADAPTER(TYPE) \
|
||||
|
|
@ -717,8 +681,9 @@ public:
|
|||
#undef ON_ADAPTER
|
||||
|
||||
void on_adapter(
|
||||
const std::string&, ngraph::ValueAccessor<std::shared_ptr<ngraph::Function>>&) override {
|
||||
// handled by `compare_functions` drop it here
|
||||
const std::string& name,
|
||||
ngraph::ValueAccessor<std::shared_ptr<ngraph::Function>>& adapter) override {
|
||||
verify_function(name, adapter);
|
||||
}
|
||||
|
||||
bool all_attr_was_compared() const {
|
||||
|
|
@ -737,6 +702,7 @@ private:
|
|||
bool should_return() const {
|
||||
return m_fast_exit && m_cmp_result.has_error();
|
||||
}
|
||||
|
||||
template <typename AttrValue>
|
||||
void verify(const std::string& name, const AttrValue& attr_value) {
|
||||
if (should_return()) {
|
||||
|
|
@ -756,23 +722,85 @@ private:
|
|||
}
|
||||
}
|
||||
|
||||
void verify_mem_buf(
|
||||
const std::string& name, const std::shared_ptr<ngraph::runtime::AlignedBuffer>& buffer) {
|
||||
if (should_return()) {
|
||||
return;
|
||||
}
|
||||
m_visited_attributes.insert(name);
|
||||
const auto ref_value = m_attr_ref.get<storage::MemoryChunk>(name);
|
||||
if (!ref_value) {
|
||||
m_cmp_result += "missing attribute name: '" + name + "'";
|
||||
return;
|
||||
}
|
||||
|
||||
if (buffer->size() != ref_value->size() ||
|
||||
std::memcmp(ref_value->data(), buffer->get_ptr(), ref_value->size()) != 0) {
|
||||
m_cmp_result += "mismatch in value: '" + name + "' : look in to the mem buffer";
|
||||
return;
|
||||
}
|
||||
}
|
||||
using FunctionAccessor = ngraph::ValueAccessor<std::shared_ptr<ngraph::Function>>;
|
||||
|
||||
void verify_function(const std::string& name, FunctionAccessor& adapter) {
|
||||
if (should_return()) {
|
||||
return;
|
||||
}
|
||||
m_visited_attributes.insert(name);
|
||||
const auto ref_value = m_attr_ref.get<std::shared_ptr<ngraph::Function>>(name);
|
||||
if (!ref_value) {
|
||||
m_cmp_result += "missing attribute name: '" + name + "'";
|
||||
return;
|
||||
}
|
||||
Comparator c(m_check_flags);
|
||||
const auto result = c.compare(*ref_value, adapter.get());
|
||||
if (!result.valid) {
|
||||
m_cmp_result += result.message;
|
||||
}
|
||||
}
|
||||
|
||||
void verify_others(const std::string& name, ngraph::ValueAccessor<void>& adapter) {
|
||||
if (auto inputs =
|
||||
ngraph::as_type<ngraph::AttributeAdapter<SubGraphOpInputDescription>>(&adapter)) {
|
||||
verify(name, inputs->get());
|
||||
} else if (
|
||||
auto outputs =
|
||||
ngraph::as_type<ngraph::AttributeAdapter<SubGraphOpOutputDescription>>(&adapter)) {
|
||||
verify(name, outputs->get());
|
||||
} else if (ngraph::is_type<ngraph::AttributeAdapter<SpecialBodyPorts>>(&adapter)) {
|
||||
// drop comparison, no more info than port indexes which will be check in
|
||||
// subgraph::compare_io
|
||||
} else if (
|
||||
auto a = ngraph::as_type<
|
||||
ngraph::AttributeAdapter<std::shared_ptr<ngraph::runtime::AlignedBuffer>>>(
|
||||
&adapter)) {
|
||||
verify_mem_buf(name, a->get());
|
||||
} else {
|
||||
m_cmp_result += "compare attr [ ERR ]: " + name +
|
||||
" [drop `void` comparison which is '" + adapter.get_type_info().name +
|
||||
"']";
|
||||
}
|
||||
}
|
||||
//-- DATA --
|
||||
const ReadAndStoreAttributes& m_attr_ref;
|
||||
Result m_cmp_result;
|
||||
Comparator::CmpValues m_check_flags;
|
||||
std::set<AttrName> m_visited_attributes;
|
||||
bool m_fast_exit{true};
|
||||
static constexpr bool m_fast_exit{true};
|
||||
};
|
||||
|
||||
} // namespace attr_comparison
|
||||
|
||||
class CompareNodesAttributes {
|
||||
public:
|
||||
CompareNodesAttributes() : m_compare_attr(m_store_attr) {}
|
||||
using ReadAndStoreAttributes = detail::ReadAndStoreAttributes;
|
||||
using ReadAndCompareAttributes = detail::ReadAndCompareAttributes;
|
||||
|
||||
attr_comparison::ReadAndStoreAttributes& get_ref_reader() {
|
||||
CompareNodesAttributes(Comparator::CmpValues m_compare_flags)
|
||||
: m_compare_attr(m_store_attr, m_compare_flags) {}
|
||||
|
||||
ReadAndStoreAttributes& get_ref_reader() {
|
||||
return m_store_attr;
|
||||
}
|
||||
|
||||
attr_comparison::ReadAndCompareAttributes& get_cmp_reader() {
|
||||
ReadAndCompareAttributes& get_cmp_reader() {
|
||||
return m_compare_attr;
|
||||
}
|
||||
|
||||
|
|
@ -787,18 +815,477 @@ public:
|
|||
}
|
||||
if (!c.m_compare_attr.all_attr_was_compared()) {
|
||||
return "not all of attr was compared: " +
|
||||
std::to_string(c.m_compare_attr.compared_attr_number()) + " vs " +
|
||||
std::to_string(c.m_store_attr.attributes_number());
|
||||
to_str(c.m_compare_attr.compared_attr_number()) + " vs " +
|
||||
to_str(c.m_store_attr.attributes_number());
|
||||
}
|
||||
return "looks good [compared " + std::to_string(c.m_compare_attr.compared_attr_number()) +
|
||||
return "looks good [compared " + to_str(c.m_compare_attr.compared_attr_number()) +
|
||||
" attributes]";
|
||||
}
|
||||
|
||||
private:
|
||||
attr_comparison::ReadAndStoreAttributes m_store_attr;
|
||||
attr_comparison::ReadAndCompareAttributes m_compare_attr;
|
||||
ReadAndStoreAttributes m_store_attr;
|
||||
ReadAndCompareAttributes m_compare_attr;
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
Comparator::Result compare(
|
||||
ngraph::Node* node1, ngraph::Node* node2, Comparator::CmpValues comparition_flags) {
|
||||
detail::CompareNodesAttributes compare_nodes_attr(comparition_flags);
|
||||
node1->visit_attributes(compare_nodes_attr.get_ref_reader());
|
||||
node2->visit_attributes(compare_nodes_attr.get_cmp_reader());
|
||||
if (!compare_nodes_attr.equal()) {
|
||||
return Comparator::Result::error(
|
||||
"Comparison of attributes failed for nodes " + name(node1) + ", " + name(node2) +
|
||||
" [cmp status: " + to_str(compare_nodes_attr) + "]");
|
||||
}
|
||||
return Comparator::Result::ok(to_str(compare_nodes_attr));
|
||||
}
|
||||
|
||||
} // namespace attributes
|
||||
|
||||
namespace subgraph {
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <typename Ptr>
|
||||
Ptr not_null(Ptr&& p) {
|
||||
if (!p) {
|
||||
THROW_IE_EXCEPTION << "empty pointer";
|
||||
}
|
||||
return std::forward<Ptr>(p);
|
||||
}
|
||||
|
||||
template <typename InOut1, typename InOut2>
|
||||
bool equal_type_and_partial_shape(const InOut1& lhs, const InOut2& rhs) {
|
||||
return lhs.get_element_type() == rhs.get_element_type() &&
|
||||
lhs.get_partial_shape() == rhs.get_partial_shape();
|
||||
}
|
||||
|
||||
class NodeAndInputDescription {
|
||||
public:
|
||||
using SubGraphOp = ngraph::op::util::SubGraphOp;
|
||||
using InputDescripton = SubGraphOp::InputDescription;
|
||||
using InputNode = ngraph::Input<ngraph::Node>;
|
||||
using Parameter = ngraph::opset6::Parameter;
|
||||
|
||||
explicit NodeAndInputDescription(
|
||||
const InputNode& input, const Parameter* parameter, const InputDescripton* description)
|
||||
: m_input(input), m_parameter(not_null(parameter)), m_description(not_null(description)) {}
|
||||
|
||||
static bool equal_descriptions(const InputDescripton* lhs, const InputDescripton* rhs) {
|
||||
if (!lhs || !rhs || lhs->get_type_info() != rhs->get_type_info()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (lhs->get_type_info() == SubGraphOp::SliceInputDescription::type_info) {
|
||||
using InDesc = SubGraphOp::SliceInputDescription;
|
||||
const InDesc* l_input = static_cast<const InDesc*>(lhs);
|
||||
const InDesc* r_input = static_cast<const InDesc*>(rhs);
|
||||
return l_input->m_start == r_input->m_start && l_input->m_stride == r_input->m_stride &&
|
||||
l_input->m_part_size == r_input->m_part_size &&
|
||||
l_input->m_end == r_input->m_end && l_input->m_axis == r_input->m_axis;
|
||||
} else if (lhs->get_type_info() == SubGraphOp::MergedInputDescription::type_info) {
|
||||
return true; // noting extra to check
|
||||
} else if (lhs->get_type_info() == SubGraphOp::InvariantInputDescription::type_info) {
|
||||
return true; // noting extra to check
|
||||
}
|
||||
|
||||
THROW_IE_EXCEPTION << "Type is not supported: [" << lhs->get_type_info().name << "]";
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
bool parameter_and_input_match(size_t num_iterations) const {
|
||||
if (const SubGraphOp::SliceInputDescription* slice_description =
|
||||
ngraph::as_type<const SubGraphOp::SliceInputDescription>(m_description)) {
|
||||
if (m_parameter->get_element_type() != m_input.get_element_type()) {
|
||||
return false;
|
||||
}
|
||||
const auto& param_partial_shape = m_parameter->get_partial_shape();
|
||||
const auto& input_partial_shape = m_input.get_partial_shape();
|
||||
if (param_partial_shape.is_dynamic() && input_partial_shape.is_dynamic()) {
|
||||
return true;
|
||||
}
|
||||
if (!param_partial_shape.is_static() || !input_partial_shape.is_static()) {
|
||||
return false;
|
||||
}
|
||||
const auto& param_shape = param_partial_shape.to_shape();
|
||||
const auto& input_shape = input_partial_shape.to_shape();
|
||||
if (param_shape.size() != input_shape.size()) {
|
||||
return false;
|
||||
}
|
||||
if (param_shape[slice_description->m_axis] != slice_description->m_part_size) {
|
||||
return false;
|
||||
}
|
||||
for (size_t i = 0; i != param_shape.size(); ++i) {
|
||||
const auto expected_axis_size =
|
||||
i == slice_description->m_axis ? slice_description->m_part_size * num_iterations
|
||||
: param_shape[i];
|
||||
if (input_shape[i] != expected_axis_size) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
} else if (
|
||||
m_description->get_type_info() == SubGraphOp::MergedInputDescription::type_info ||
|
||||
m_description->get_type_info() == SubGraphOp::InvariantInputDescription::type_info) {
|
||||
return equal_type_and_partial_shape(*m_parameter, m_input);
|
||||
}
|
||||
|
||||
THROW_IE_EXCEPTION << "Type is not supported: [" << m_description->get_type_info().name
|
||||
<< "]";
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
static bool equal_parameters(const Parameter* lhs, const Parameter* rhs) {
|
||||
return lhs && rhs && equal_type_and_partial_shape(*lhs, *rhs);
|
||||
}
|
||||
|
||||
friend bool operator==(const NodeAndInputDescription& lhs, const NodeAndInputDescription& rhs) {
|
||||
if (!equal_descriptions(lhs.m_description, rhs.m_description)) {
|
||||
return false;
|
||||
}
|
||||
return equal_parameters(lhs.m_parameter, rhs.m_parameter);
|
||||
}
|
||||
|
||||
private:
|
||||
const InputNode m_input;
|
||||
const Parameter* m_parameter;
|
||||
const InputDescripton* m_description;
|
||||
};
|
||||
|
||||
class NodeAndOutputDescription {
|
||||
public:
|
||||
using SubGraphOp = ngraph::op::util::SubGraphOp;
|
||||
using OutputDescription = SubGraphOp::OutputDescription;
|
||||
using OutputNode = ngraph::Output<ngraph::Node>;
|
||||
using Result = ngraph::opset6::Result;
|
||||
|
||||
explicit NodeAndOutputDescription(
|
||||
const OutputNode& output, const Result* result, const OutputDescription* description)
|
||||
: m_output(output), m_result(not_null(result)), m_description(not_null(description)) {}
|
||||
|
||||
static bool equal_descriptions(const OutputDescription* lhs, const OutputDescription* rhs) {
|
||||
if (!lhs || !rhs || lhs->get_type_info() != rhs->get_type_info()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (lhs->get_type_info() == SubGraphOp::ConcatOutputDescription::type_info) {
|
||||
using OutDesc = SubGraphOp::ConcatOutputDescription;
|
||||
const OutDesc* l_output = static_cast<const OutDesc*>(lhs);
|
||||
const OutDesc* r_output = static_cast<const OutDesc*>(rhs);
|
||||
return l_output->m_start == r_output->m_start &&
|
||||
l_output->m_stride == r_output->m_stride &&
|
||||
l_output->m_part_size == r_output->m_part_size &&
|
||||
l_output->m_end == r_output->m_end && l_output->m_axis == r_output->m_axis;
|
||||
} else if (lhs->get_type_info() == SubGraphOp::BodyOutputDescription::type_info) {
|
||||
using OutDesc = SubGraphOp::BodyOutputDescription;
|
||||
const OutDesc* l_output = static_cast<const OutDesc*>(lhs);
|
||||
const OutDesc* r_output = static_cast<const OutDesc*>(rhs);
|
||||
return l_output->m_iteration == r_output->m_iteration;
|
||||
}
|
||||
|
||||
THROW_IE_EXCEPTION << "Type is not supported: [" << lhs->get_type_info().name << "]";
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
bool result_and_output_match(size_t num_iterations) const {
|
||||
if (const auto concat_desciption =
|
||||
ngraph::as_type<const SubGraphOp::ConcatOutputDescription>(m_description)) {
|
||||
if (m_result->output(0).get_element_type() != m_output.get_element_type()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const auto& output_partial_shape = m_output.get_partial_shape();
|
||||
const auto& result_partial_shape = m_result->output(0).get_partial_shape();
|
||||
if (result_partial_shape.is_dynamic() && output_partial_shape.is_dynamic()) {
|
||||
return true;
|
||||
}
|
||||
if (!result_partial_shape.is_static() || !output_partial_shape.is_static()) {
|
||||
return false;
|
||||
}
|
||||
const auto& output_shape = output_partial_shape.to_shape();
|
||||
const auto& result_shape = result_partial_shape.to_shape();
|
||||
if (result_shape.size() != output_shape.size()) {
|
||||
return false;
|
||||
}
|
||||
for (size_t i = 0; i != result_shape.size(); ++i) {
|
||||
const auto axis_multiplier = i == concat_desciption->m_axis ? num_iterations : 1;
|
||||
if (result_shape[i] * axis_multiplier != output_shape[i]) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
} else if (m_description->get_type_info() == SubGraphOp::BodyOutputDescription::type_info) {
|
||||
return equal_type_and_partial_shape(m_result->output(0), m_output);
|
||||
}
|
||||
|
||||
THROW_IE_EXCEPTION << "Type is not supported: [" << m_description->get_type_info().name
|
||||
<< "]";
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
static bool equal_results(const Result* lhs, const Result* rhs) {
|
||||
return lhs && rhs && equal_type_and_partial_shape(lhs->output(0), rhs->output(0));
|
||||
}
|
||||
|
||||
friend bool operator==(
|
||||
const NodeAndOutputDescription& lhs, const NodeAndOutputDescription& rhs) {
|
||||
if (!equal_descriptions(lhs.m_description, rhs.m_description)) {
|
||||
return false;
|
||||
}
|
||||
return equal_results(lhs.m_result, rhs.m_result);
|
||||
}
|
||||
|
||||
private:
|
||||
const OutputNode m_output;
|
||||
const Result* m_result;
|
||||
const OutputDescription* m_description;
|
||||
};
|
||||
|
||||
class BackEdge {
|
||||
public:
|
||||
using Parameter = ngraph::opset6::Parameter;
|
||||
using Result = ngraph::opset6::Result;
|
||||
using Id = uint64_t;
|
||||
|
||||
explicit BackEdge(const Parameter* parameter, const Result* result)
|
||||
: m_parameter(not_null(parameter)), m_result(not_null(result)) {}
|
||||
|
||||
bool result_and_parameter_match() const {
|
||||
return equal_type_and_partial_shape(m_result->output(0), *m_parameter);
|
||||
}
|
||||
|
||||
friend bool operator==(const BackEdge& lhs, const BackEdge& rhs) {
|
||||
return equal_type_and_partial_shape(*lhs.m_parameter, *rhs.m_parameter) &&
|
||||
equal_type_and_partial_shape(lhs.m_result->output(0), rhs.m_result->output(0));
|
||||
}
|
||||
|
||||
private:
|
||||
const Parameter* m_parameter;
|
||||
const Result* m_result;
|
||||
};
|
||||
|
||||
std::vector<NodeAndInputDescription> extract_inputs(ngraph::op::util::SubGraphOp* sub) {
|
||||
std::vector<NodeAndInputDescription> nodes;
|
||||
const auto& fn_body = sub->get_function();
|
||||
const auto& fn_parameters = fn_body->get_parameters();
|
||||
|
||||
for (const auto& in_desc : sub->get_input_descriptions()) {
|
||||
const auto parameter = fn_parameters.at(in_desc->m_body_parameter_index).get();
|
||||
const auto input = sub->input(in_desc->m_input_index);
|
||||
nodes.push_back(NodeAndInputDescription{input, parameter, in_desc.get()});
|
||||
}
|
||||
return nodes;
|
||||
}
|
||||
|
||||
std::vector<NodeAndOutputDescription> extract_outputs(ngraph::op::util::SubGraphOp* sub) {
|
||||
std::vector<NodeAndOutputDescription> nodes;
|
||||
const auto& fn_body = sub->get_function();
|
||||
const auto& fs_results = fn_body->get_results();
|
||||
|
||||
for (const auto& out_desc : sub->get_output_descriptions()) {
|
||||
const auto result = fs_results.at(out_desc->m_body_value_index).get();
|
||||
const auto output = sub->output(out_desc->m_output_index);
|
||||
nodes.push_back(NodeAndOutputDescription{output, result, out_desc.get()});
|
||||
}
|
||||
return nodes;
|
||||
}
|
||||
|
||||
std::vector<BackEdge> extract_backedges(ngraph::op::util::SubGraphOp* sub) {
|
||||
using MergedInputDescription = ngraph::op::util::SubGraphOp::MergedInputDescription;
|
||||
std::vector<BackEdge> edges;
|
||||
const auto& fn_body = sub->get_function();
|
||||
|
||||
const auto& fs_parameters = fn_body->get_parameters();
|
||||
const auto& fs_results = fn_body->get_results();
|
||||
|
||||
for (const auto& in_desc : sub->get_input_descriptions()) {
|
||||
if (const auto& merged_in_desc =
|
||||
ngraph::as_type_ptr<const MergedInputDescription>(in_desc)) {
|
||||
const auto parameter = fs_parameters.at(merged_in_desc->m_body_parameter_index);
|
||||
const auto result = fs_results.at(merged_in_desc->m_body_value_index);
|
||||
edges.push_back(BackEdge{parameter.get(), result.get()});
|
||||
}
|
||||
}
|
||||
return edges;
|
||||
}
|
||||
|
||||
struct NotValidInputOrOutput {
|
||||
NotValidInputOrOutput(int64_t num_iterations) : m_num_iterations(num_iterations) {}
|
||||
|
||||
bool operator()(const NodeAndOutputDescription& d) const {
|
||||
return !d.result_and_output_match(m_num_iterations);
|
||||
}
|
||||
|
||||
bool operator()(const NodeAndInputDescription& d) const {
|
||||
return !d.parameter_and_input_match(m_num_iterations);
|
||||
}
|
||||
|
||||
int64_t m_num_iterations;
|
||||
};
|
||||
|
||||
bool not_valid_back_edge(const BackEdge& be) {
|
||||
return !be.result_and_parameter_match();
|
||||
}
|
||||
|
||||
bool equal_body_ports(ngraph::opset6::Loop* lhs, ngraph::opset6::Loop* rhs) {
|
||||
if (!lhs || !rhs) {
|
||||
return false;
|
||||
}
|
||||
const auto& lhs_fn_body = lhs->get_function();
|
||||
const auto& rhs_fn_body = rhs->get_function();
|
||||
|
||||
const auto& lhs_sbp = lhs->get_special_body_ports();
|
||||
const auto& rhs_sbp = rhs->get_special_body_ports();
|
||||
|
||||
constexpr int64_t port_not_provided = -1;
|
||||
|
||||
const bool input_provided = lhs_sbp.current_iteration_input_idx != port_not_provided ||
|
||||
rhs_sbp.current_iteration_input_idx != port_not_provided;
|
||||
|
||||
if (input_provided) {
|
||||
const auto& lhs_parameter =
|
||||
lhs_fn_body->get_parameters().at(lhs_sbp.current_iteration_input_idx);
|
||||
const auto& rhs_parameter =
|
||||
rhs_fn_body->get_parameters().at(rhs_sbp.current_iteration_input_idx);
|
||||
if (!NodeAndInputDescription::equal_parameters(lhs_parameter.get(), rhs_parameter.get())) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
const auto& lhs_result = lhs_fn_body->get_results().at(lhs_sbp.body_condition_output_idx);
|
||||
const auto& rhs_result = rhs_fn_body->get_results().at(rhs_sbp.body_condition_output_idx);
|
||||
|
||||
return NodeAndOutputDescription::equal_results(lhs_result.get(), rhs_result.get());
|
||||
}
|
||||
|
||||
class CompareSubGraphs {
|
||||
public:
|
||||
using Result = Comparator::Result;
|
||||
using SubGraphOp = ngraph::op::util::SubGraphOp;
|
||||
|
||||
Result compare(SubGraphOp* sub_lhs, SubGraphOp* sub_rhs) {
|
||||
const auto lhs_it_no = get_num_iterations(sub_lhs);
|
||||
const auto rhs_it_no = get_num_iterations(sub_rhs);
|
||||
if (lhs_it_no != rhs_it_no) {
|
||||
return Result::error("different number of iterations");
|
||||
}
|
||||
|
||||
not_valid_input_output = lhs_it_no;
|
||||
|
||||
const auto result_for_inputs = compare_inputs(sub_lhs, sub_rhs);
|
||||
if (!result_for_inputs.valid) {
|
||||
return result_for_inputs;
|
||||
}
|
||||
|
||||
const auto result_for_outputs = compare_outputs(sub_lhs, sub_rhs);
|
||||
if (!result_for_outputs.valid) {
|
||||
return result_for_outputs;
|
||||
}
|
||||
|
||||
return compare_backedges(sub_lhs, sub_rhs);
|
||||
}
|
||||
|
||||
private:
|
||||
Result compare_inputs(SubGraphOp* sub_lhs, SubGraphOp* sub_rhs) const {
|
||||
const auto& lhs_sub_inputs = extract_inputs(sub_lhs);
|
||||
const auto& rhs_sub_inputs = extract_inputs(sub_rhs);
|
||||
|
||||
if (lhs_sub_inputs.empty() || rhs_sub_inputs.empty()) {
|
||||
return Result::error("no input in subgraph");
|
||||
}
|
||||
|
||||
if (std::any_of(begin(lhs_sub_inputs), end(lhs_sub_inputs), not_valid_input_output)) {
|
||||
return Result::error("inputs and parameters mismatch");
|
||||
}
|
||||
if (std::any_of(begin(rhs_sub_inputs), end(rhs_sub_inputs), not_valid_input_output)) {
|
||||
return Result::error("inputs and parameters mismatch");
|
||||
}
|
||||
|
||||
if (lhs_sub_inputs.size() != rhs_sub_inputs.size() ||
|
||||
!std::is_permutation(
|
||||
begin(lhs_sub_inputs), end(lhs_sub_inputs), begin(rhs_sub_inputs))) {
|
||||
return Result::error("different SubGraph InputDescription");
|
||||
}
|
||||
return Result::ok();
|
||||
}
|
||||
|
||||
Result compare_outputs(SubGraphOp* sub_lhs, SubGraphOp* sub_rhs) const {
|
||||
const auto& lhs_sub_outputs = extract_outputs(sub_lhs);
|
||||
const auto& rhs_sub_outputs = extract_outputs(sub_rhs);
|
||||
|
||||
if (lhs_sub_outputs.empty() || rhs_sub_outputs.empty()) {
|
||||
return Result::error("no output in subgraph");
|
||||
}
|
||||
|
||||
if (std::any_of(begin(lhs_sub_outputs), end(lhs_sub_outputs), not_valid_input_output)) {
|
||||
return Result::error("outputs and results mismatch");
|
||||
}
|
||||
if (std::any_of(begin(rhs_sub_outputs), end(rhs_sub_outputs), not_valid_input_output)) {
|
||||
return Result::error("outputs and results mismatch");
|
||||
}
|
||||
|
||||
if (lhs_sub_outputs.size() != rhs_sub_outputs.size() ||
|
||||
!std::is_permutation(
|
||||
begin(lhs_sub_outputs), end(lhs_sub_outputs), begin(rhs_sub_outputs))) {
|
||||
return Result::error("different SubGraph OutputDescription");
|
||||
}
|
||||
return Result::ok();
|
||||
}
|
||||
|
||||
Result compare_backedges(SubGraphOp* sub_lhs, SubGraphOp* sub_rhs) const {
|
||||
const auto lhs_back_edges = extract_backedges(sub_lhs);
|
||||
const auto rhs_back_edges = extract_backedges(sub_rhs);
|
||||
|
||||
if (std::any_of(begin(lhs_back_edges), end(lhs_back_edges), not_valid_back_edge)) {
|
||||
return Result::error("back edges mismatch");
|
||||
}
|
||||
if (std::any_of(begin(rhs_back_edges), end(rhs_back_edges), not_valid_back_edge)) {
|
||||
return Result::error("back edges mismatch");
|
||||
}
|
||||
|
||||
if (lhs_back_edges.size() != rhs_back_edges.size() ||
|
||||
!std::is_permutation(
|
||||
begin(lhs_back_edges), end(lhs_back_edges), begin(rhs_back_edges))) {
|
||||
return Result::error("different SubGraph BackEdges");
|
||||
}
|
||||
if (auto loop_lhs = ngraph::as_type<ngraph::opset6::Loop>(sub_lhs)) {
|
||||
auto loop_rhs = ngraph::as_type<ngraph::opset6::Loop>(sub_rhs);
|
||||
if (!equal_body_ports(loop_lhs, loop_rhs)) {
|
||||
return Result::error("different Special Body Ports");
|
||||
}
|
||||
}
|
||||
return Result::ok();
|
||||
}
|
||||
|
||||
static int64_t get_num_iterations(ngraph::op::util::SubGraphOp* sub) {
|
||||
using namespace ngraph::opset6;
|
||||
if (const auto ti = dynamic_cast<const TensorIterator*>(sub)) {
|
||||
return ti->get_num_iterations();
|
||||
}
|
||||
if (const auto l = dynamic_cast<const Loop*>(sub)) {
|
||||
return l->get_num_iterations();
|
||||
}
|
||||
|
||||
return -1;
|
||||
}
|
||||
|
||||
NotValidInputOrOutput not_valid_input_output{-1};
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
Comparator::Result compare_io(
|
||||
ngraph::op::util::SubGraphOp* sub_lhs, ngraph::op::util::SubGraphOp* sub_rhs) {
|
||||
return detail::CompareSubGraphs{}.compare(sub_lhs, sub_rhs);
|
||||
}
|
||||
} // namespace subgraph
|
||||
|
||||
Comparator::Result Comparator::compare(
|
||||
const std::shared_ptr<ngraph::Function>& f1, const std::shared_ptr<ngraph::Function>& f2) {
|
||||
/*
|
||||
|
|
@ -874,7 +1361,7 @@ Comparator::Result Comparator::compare(
|
|||
auto subgraph2 = dynamic_cast<ngraph::op::util::SubGraphOp*>(node2);
|
||||
|
||||
if (subgraph1 && subgraph2) {
|
||||
auto result = recreate().compare(subgraph1->get_function(), subgraph2->get_function());
|
||||
const auto result = subgraph::compare_io(subgraph1, subgraph2);
|
||||
if (!result.valid) {
|
||||
return result;
|
||||
}
|
||||
|
|
@ -901,14 +1388,15 @@ Comparator::Result Comparator::compare(
|
|||
name(node1) + " and " + to_str(node2->inputs().size()) + " for " + name(node2));
|
||||
}
|
||||
|
||||
for (int i = 0; i < node1->inputs().size(); ++i) {
|
||||
for (size_t i = 0; i < node1->inputs().size(); ++i) {
|
||||
if (should_compare(CmpValues::CONST_VALUES)) {
|
||||
using Constant = ngraph::opset1::Constant;
|
||||
const auto equal_value =
|
||||
::attributes::detail::equal::Equal<std::shared_ptr<Constant>>::equal_value;
|
||||
|
||||
auto const1 = ngraph::as_type_ptr<Constant>(node1->get_input_node_shared_ptr(i));
|
||||
auto const2 = ngraph::as_type_ptr<Constant>(node2->get_input_node_shared_ptr(i));
|
||||
using namespace ::attr_comparison::equal;
|
||||
if (const1 && const2 &&
|
||||
!Equal<std::shared_ptr<Constant>>::equal_value(const1, const2)) {
|
||||
if (const1 && const2 && !equal_value(const1, const2)) {
|
||||
err_log << "Different Constant values detected\n"
|
||||
<< node1->description() << " Input(" << i << ") and "
|
||||
<< node2->description() << " Input(" << i << ")" << std::endl;
|
||||
|
|
@ -969,13 +1457,9 @@ Comparator::Result Comparator::compare(
|
|||
}
|
||||
|
||||
if (should_compare(CmpValues::ATTRIBUTES)) {
|
||||
CompareNodesAttributes compare_nodes;
|
||||
node1->visit_attributes(compare_nodes.get_ref_reader());
|
||||
node2->visit_attributes(compare_nodes.get_cmp_reader());
|
||||
if (!compare_nodes.equal()) {
|
||||
return Result::error(
|
||||
"Comparison of attributes failed for nodes " + name(node1) + ", " + name(node2) +
|
||||
" [cmp status: " + to_string(compare_nodes) + "]");
|
||||
const auto result = attributes::compare(node1, node2, m_comparition_flags);
|
||||
if (!result.valid) {
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -990,7 +1474,6 @@ void Comparator::add_nodes_inputs_to_queue(ngraph::Node* node1, ngraph::Node* no
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
FunctionsComparator::Result FunctionsComparator::compare(
|
||||
|
|
|
|||
|
|
@ -71,7 +71,7 @@ private:
|
|||
|
||||
///
|
||||
/// \deprecated
|
||||
/// \brief compare_functions is obsolete function use FunctionComparator instead.
|
||||
/// \brief compare_functions is obsolete function use FunctionsComparator instead.
|
||||
///
|
||||
inline std::pair<bool, std::string> compare_functions(
|
||||
const std::shared_ptr<ngraph::Function>& f1,
|
||||
|
|
|
|||
|
|
@ -51,12 +51,9 @@ namespace ngraph
|
|||
std::shared_ptr<Function> get_function() override;
|
||||
|
||||
int64_t get_num_iterations() const { return m_num_iterations; }
|
||||
void set_num_iterations(int64_t num_iterations)
|
||||
{
|
||||
m_num_iterations = num_iterations;
|
||||
}
|
||||
|
||||
private:
|
||||
void try_to_set_num_iterations_if_no_slice_inputs();
|
||||
|
||||
int64_t m_num_iterations = -1;
|
||||
};
|
||||
}
|
||||
|
|
|
|||
|
|
@ -37,14 +37,6 @@ bool op::v0::TensorIterator::visit_attributes(AttributeVisitor& visitor)
|
|||
visitor.on_attribute("input_descriptions", m_input_descriptions);
|
||||
visitor.on_attribute("output_descriptions", m_output_descriptions);
|
||||
|
||||
for (const auto& output_description : m_output_descriptions)
|
||||
{
|
||||
if (auto concat = as_type_ptr<ConcatOutputDescription>(output_description))
|
||||
{
|
||||
m_num_iterations = ((std::abs(concat->m_end - concat->m_start)) / concat->m_part_size);
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
|
|
@ -176,6 +168,8 @@ void op::v0::TensorIterator::validate_and_infer_types()
|
|||
revalidate_and_infer_types_for_body_ops();
|
||||
|
||||
// Output
|
||||
try_to_set_num_iterations_if_no_slice_inputs();
|
||||
|
||||
index_it = 0;
|
||||
for (const auto& output_description : m_output_descriptions)
|
||||
{
|
||||
|
|
@ -241,6 +235,35 @@ std::shared_ptr<Function> op::v0::TensorIterator::get_function()
|
|||
return get_body();
|
||||
}
|
||||
|
||||
namespace
|
||||
{
|
||||
template <typename Desc>
|
||||
bool has_slice_input_desc(const Desc& desc)
|
||||
{
|
||||
const auto is_slice_input_desc = +[](typename Desc::const_reference d) {
|
||||
return is_type<op::util::SubGraphOp::SliceInputDescription>(d);
|
||||
};
|
||||
return std::any_of(begin(desc), end(desc), is_slice_input_desc);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
void op::v0::TensorIterator::try_to_set_num_iterations_if_no_slice_inputs()
|
||||
{
|
||||
if (m_num_iterations != -1 || has_slice_input_desc(get_input_descriptions()))
|
||||
{
|
||||
return;
|
||||
}
|
||||
|
||||
for (const auto& output_description : m_output_descriptions)
|
||||
{
|
||||
if (auto concat = as_type_ptr<ConcatOutputDescription>(output_description))
|
||||
{
|
||||
m_num_iterations = ((std::abs(concat->m_end - concat->m_start)) / concat->m_part_size);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
std::shared_ptr<Node>
|
||||
op::v0::TensorIterator::clone_with_new_inputs(const OutputVector& new_args) const
|
||||
{
|
||||
|
|
|
|||
Loading…
Reference in New Issue