Merge TensorFlow Frontend Implementation (#7855)

* Migrate POC for TensorFlow frontend

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Refactor InputModelTensorFlow API

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Repack POC to official API

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Remove tensorflow API from public include

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Make TF frontend work from MO and clean-up code

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Apply codestyle

* Fix win biuld

* Fix Linux build

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Implement Place class

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Determine outputs from graph

* Implement all Place classes

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Make small clean-up

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Apply code-style corrections

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Determine cut nodes

* Apply codestyle

* Rework to use places

* Fix conversion issue

* Fix build

* Fix conversion

* Small fixes

* Add test for tf frontend

* Add tests

* Implement partial conversion

* Use dynamic type in TFFrameworkNode

* Fix build on Linux

* Implement InputModelTF class

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Fix code by replacing InputModelTensorFlow to InputModelTF

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Fix to pass getPlaceByTensorName test

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Refactor and clean the code

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Finalize refactoring code

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Support freezing inputs

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Add support for pruning input ports as new model output

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Apply code-style fixes

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* move op convertors to separate files, refactoring

* openvino codestyle

* openvino codestyle

* fix crash of layer tests

* fix missprint

* Implement TensorFlow NodeContext and DecoderTFProto classes

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Switch to new NodeContext

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Remove ngraph_builder class and node_context_impl class

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Move decoder/graph_iterator to separate files and remove old files

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Document Decoder, GraphIterator, and NodeContext classes

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Apply code style

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Remove empty file graph_iterator_proto.cpp and redundant comments

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Use base class for GraphIterator in model class and correct exception class

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Use ends_with from util library

* Remain only InputModelTF constructor with GraphIterator and adopt other code

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Correct code after merge

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Apply code style

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Fix code based on feedback: delete extra namespace usage, etc.

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Make legacy TF frontend default in MO and avoid reading tf models from ReadNetwork

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Fix build issue with FrameworkNode after upstream merge

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Correct frontend name in test and clean the code

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* refactoring of tf FrontEnd: rename namespaces, delete default opset

* codestyle

* fix e2e tests

* change namespaces of external classes

* resolve review comment

* codestyle

* Clean useless comments and fix tests

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Add copyright for proto files

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Make op translators a part of FrontEndTF

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Align CMakeLists.txt file with other frontends

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Add job_pool into CMakeLists.txt for tensorflow tests

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Remove cout and WA in ie_network_reader.cpp

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Correct DecoderBase and FrontEndTF methods, clean the code from if 0, utility.hpp, ngraph_conversions.hpp

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Clean operation translator code and remove empty implementation

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Extend cmake template files with tensorflow frontend

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Apply code-style

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Fix build issues on MacOS and apply review feedback

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Remove ConstantFolding pass, add destructors for DecoderBase, GraphIterator, and fix Win build

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Fix Windows build issue

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Apply code-style fix

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Fix issues with code style: remove virtual keyword, OpMap removal

Signed-off-by: Roman Kazantsev <roman.kazantsev@intel.com>

* Resolve review remarks

* Resolve review remarks

* codestyle

* resolve review remarks

* fix copyright

* resolve review comments

Co-authored-by: Maxim Vafin <maxim.vafin@intel.com>
Co-authored-by: Ivan Tikhonov <ivan.tikhonov@intel.com>
This commit is contained in:
Roman Kazantsev 2021-10-14 10:36:10 +03:00 committed by GitHub
parent a8f0109584
commit 1752298e17
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23
105 changed files with 7481 additions and 3 deletions

View File

@ -123,10 +123,11 @@ endif()
ie_dependent_option(NGRAPH_ONNX_FRONTEND_ENABLE "Enable ONNX FrontEnd" ON "protoc_available" OFF)
ie_dependent_option(NGRAPH_PDPD_FRONTEND_ENABLE "Enable PaddlePaddle FrontEnd" ON "protoc_available" OFF)
ie_option(NGRAPH_IR_FRONTEND_ENABLE "Enable IR FrontEnd" ON)
ie_dependent_option(NGRAPH_TF_FRONTEND_ENABLE "Enable TensorFlow FrontEnd" ON "protoc_available" OFF)
ie_dependent_option(NGRAPH_USE_PROTOBUF_LITE "Compiles and links with protobuf-lite" ON
"NGRAPH_ONNX_FRONTEND_ENABLE" OFF)
ie_dependent_option(NGRAPH_USE_SYSTEM_PROTOBUF "Use system protobuf" OFF
"NGRAPH_ONNX_FRONTEND_ENABLE OR NGRAPH_PDPD_FRONTEND_ENABLE" OFF)
"NGRAPH_ONNX_FRONTEND_ENABLE OR NGRAPH_PDPD_FRONTEND_ENABLE OR NGRAPH_TF_FRONTEND_ENABLE" OFF)
ie_dependent_option(NGRAPH_UNIT_TEST_ENABLE "Enables ngraph unit tests" ON "ENABLE_TESTS;NOT ANDROID" OFF)
ie_dependent_option(NGRAPH_UNIT_TEST_BACKENDS_ENABLE "Control the building of unit tests using backends" ON
"NGRAPH_UNIT_TEST_ENABLE" OFF)

View File

@ -169,9 +169,11 @@ set(${CMAKE_FIND_PACKAGE_NAME}_Runtime_FOUND ON)
set(${CMAKE_FIND_PACKAGE_NAME}_ONNX_FOUND @NGRAPH_ONNX_FRONTEND_ENABLE@)
set(${CMAKE_FIND_PACKAGE_NAME}_PaddlePaddle_FOUND @NGRAPH_PDPD_FRONTEND_ENABLE@)
set(${CMAKE_FIND_PACKAGE_NAME}_TensorFlow_FOUND @NGRAPH_TF_FRONTEND_ENABLE@)
set(${CMAKE_FIND_PACKAGE_NAME}_Frontend_ONNX_FOUND ${${CMAKE_FIND_PACKAGE_NAME}_ONNX_FOUND})
set(${CMAKE_FIND_PACKAGE_NAME}_Frontend_PaddlePaddle_FOUND ${${CMAKE_FIND_PACKAGE_NAME}_PaddlePaddle_FOUND})
set(${CMAKE_FIND_PACKAGE_NAME}_Frontend_TensorFlow_FOUND ${${CMAKE_FIND_PACKAGE_NAME}_TensorFlow_FOUND})
set(${CMAKE_FIND_PACKAGE_NAME}_Frontend_IR_FOUND ${${CMAKE_FIND_PACKAGE_NAME}_IR_FOUND})
# if no components specified, only Runtime is provided
@ -185,7 +187,7 @@ endif()
foreach(target openvino::runtime openvino::runtime::c openvino::core
openvino::frontend::manager openvino::frontend::onnx
openvino::frontend::paddlepaddle)
openvino::frontend::paddlepaddle openvino::frontend::tensorflow)
if(TARGET ${target} AND _ov_as_external_package)
_ov_target_no_deprecation_error(${target})
endif()
@ -205,3 +207,4 @@ endif()
unset(${CMAKE_FIND_PACKAGE_NAME}_PaddlePaddle_FOUND)
unset(${CMAKE_FIND_PACKAGE_NAME}_ONNX_FOUND)
unset(${CMAKE_FIND_PACKAGE_NAME}_TensorFlow_FOUND)

View File

@ -70,6 +70,12 @@ if(TARGET openvino::frontend::paddlepaddle AND NOT TARGET ngraph::paddlepaddle_n
INTERFACE_LINK_LIBRARIES openvino::frontend::paddlepaddle)
endif()
if(TARGET openvino::frontend::tensorflow AND NOT TARGET ngraph::tensorflow_ngraph_frontend)
add_library(ngraph::tensorflow_ngraph_frontend INTERFACE IMPORTED)
set_target_properties(ngraph::tensorflow_ngraph_frontend PROPERTIES
INTERFACE_LINK_LIBRARIES openvino::frontend::tensorflow)
endif()
set(ngraph_ngraph_FOUND ON)
set(NGRAPH_LIBRARIES ngraph::ngraph)
@ -88,6 +94,7 @@ if(ngraph_onnx_importer_FOUND)
endif()
set(ngraph_paddlepaddle_frontend_FOUND ${OpenVINO_Frontend_PaddlePaddle_FOUND})
set(ngraph_tensorflow_frontend_FOUND ${OpenVINO_Frontend_TensorFlow_FOUND})
set(ngraph_ir_frontend_FOUND ${OpenVINO_Frontend_IR_FOUND})
check_required_components(ngraph)

View File

@ -121,6 +121,7 @@ def get_moc_frontends(argv: argparse.Namespace):
# Set which frontend to use by default, values should be 'new' or 'legacy'
frontend_defaults = {
'onnx': 'legacy',
'tf': 'legacy'
}
# Disable MOC frontend if default is set to legacy and no user override
if frontend_defaults.get(moc_front_end.get_name()) == 'legacy' and not use_new_frontend:

View File

@ -17,3 +17,7 @@ endif()
if (NGRAPH_IR_FRONTEND_ENABLE)
add_subdirectory(ir)
endif()
if (NGRAPH_TF_FRONTEND_ENABLE)
add_subdirectory(tensorflow)
endif()

View File

@ -0,0 +1,99 @@
# Copyright (C) 2018-2021 Intel Corporation
# SPDX-License-Identifier: Apache-2.0
#
set(TARGET_NAME "tensorflow_ngraph_frontend")
file(GLOB_RECURSE LIBRARY_SRC ${CMAKE_CURRENT_SOURCE_DIR}/src/*.cpp)
file(GLOB_RECURSE LIBRARY_HEADERS ${CMAKE_CURRENT_SOURCE_DIR}/src/*.hpp)
file(GLOB_RECURSE LIBRARY_PUBLIC_HEADERS ${CMAKE_CURRENT_SOURCE_DIR}/include/*.hpp)
set(${TARGET_NAME}_INCLUDE_DIR ${CMAKE_CURRENT_SOURCE_DIR}/include)
# Create named folders for the sources within the .vcproj
# Empty name lists them directly under the .vcproj
source_group("src" FILES ${LIBRARY_SRC})
source_group("include" FILES ${LIBRARY_HEADERS})
source_group("public include" FILES ${LIBRARY_PUBLIC_HEADERS})
set(PROTO_SRCS)
set(PROTO_HDRS)
# Generate protobuf file on build time for each '.proto' file in src/proto
file(GLOB proto_files ${CMAKE_CURRENT_SOURCE_DIR}/src/proto/*.proto)
foreach(INFILE ${proto_files})
get_filename_component(FILE_DIR ${INFILE} DIRECTORY)
get_filename_component(FILE_WE ${INFILE} NAME_WE)
set(OUTPUT_PB_SRC ${CMAKE_CURRENT_BINARY_DIR}/${FILE_WE}.pb.cc)
set(OUTPUT_PB_HEADER ${CMAKE_CURRENT_BINARY_DIR}/${FILE_WE}.pb.h)
set(GENERATED_PROTO ${INFILE})
add_custom_command(
OUTPUT "${OUTPUT_PB_SRC}" "${OUTPUT_PB_HEADER}"
COMMAND ${PROTOC_EXECUTABLE} ARGS --cpp_out ${CMAKE_CURRENT_BINARY_DIR} -I ${FILE_DIR} ${FILE_WE}.proto
DEPENDS ${PROTOC_EXECUTABLE} ${GENERATED_PROTO}
COMMENT "Running C++ protocol buffer compiler (${PROTOC_EXECUTABLE}) on ${GENERATED_PROTO}"
VERBATIM
COMMAND_EXPAND_LISTS)
list(APPEND PROTO_SRCS "${OUTPUT_PB_SRC}")
list(APPEND PROTO_HDRS "${OUTPUT_PB_HEADER}")
endforeach()
add_custom_target(${TARGET_NAME}_proto DEPENDS ${PROTO_SRCS} ${PROTO_HDRS})
set_source_files_properties(${PROTO_SRCS} ${PROTO_HDRS} PROPERTIES GENERATED TRUE)
# Disable all warnings for generated code
set_source_files_properties(${PROTO_SRCS} ${PROTO_HDRS} PROPERTIES COMPILE_OPTIONS -w)
# Create shared library
add_library(${TARGET_NAME} SHARED ${LIBRARY_SRC} ${LIBRARY_HEADERS} ${LIBRARY_PUBLIC_HEADERS} ${PROTO_SRCS} ${PROTO_HDRS})
add_library(openvino::frontend::tensorflow ALIAS ${TARGET_NAME})
add_dependencies(${TARGET_NAME} tensorflow_ngraph_frontend_proto)
ov_ncc_naming_style(FOR_TARGET ${TARGET_NAME}
INCLUDE_DIRECTORY "${CMAKE_CURRENT_SOURCE_DIR}/include"
ADDITIONAL_INCLUDE_DIRECTORIES
$<TARGET_PROPERTY:frontend_manager::static,INTERFACE_INCLUDE_DIRECTORIES>)
target_include_directories(${TARGET_NAME}
PUBLIC
$<BUILD_INTERFACE:${${TARGET_NAME}_INCLUDE_DIR}>
$<INSTALL_INTERFACE:${FRONTEND_INSTALL_INCLUDE}>
PRIVATE
${CMAKE_CURRENT_SOURCE_DIR}/src
${CMAKE_CURRENT_BINARY_DIR})
target_include_directories(${TARGET_NAME} SYSTEM PRIVATE ${Protobuf_INCLUDE_DIRS}
${CMAKE_CURRENT_BINARY_DIR})
if(COMMAND ie_add_vs_version_file)
ie_add_vs_version_file(NAME ${TARGET_NAME}
FILEDESCRIPTION "FrontEnd to load and convert TensorFlow file format")
endif()
link_system_libraries(${TARGET_NAME} PRIVATE ${Protobuf_LITE_LIBRARIES})
target_link_libraries(${TARGET_NAME} PRIVATE frontend_manager::static
PRIVATE ngraph::builder inference_engine_transformations libprotobuf openvino::util)
add_clang_format_target(${TARGET_NAME}_clang FOR_TARGETS ${TARGET_NAME}
EXCLUDE_PATTERNS ${PROTO_SRCS} ${PROTO_HDRS})
set_target_properties(${TARGET_NAME} PROPERTIES EXPORT_NAME frontend::tensorflow)
# TODO: add install commands once TensorFlow frontend is complete
#install(TARGETS ${TARGET_NAME} EXPORT OpenVINOTargets
# RUNTIME DESTINATION ${IE_CPACK_RUNTIME_PATH} COMPONENT ngraph
# ARCHIVE DESTINATION ${IE_CPACK_ARCHIVE_PATH} COMPONENT ngraph
# LIBRARY DESTINATION ${IE_CPACK_LIBRARY_PATH} COMPONENT ngraph)
#install(DIRECTORY ${${TARGET_NAME}_INCLUDE_DIR}/tensorflow_frontend
# DESTINATION ${FRONTEND_INSTALL_INCLUDE}
# COMPONENT ngraph_dev
# FILES_MATCHING PATTERN "*.hpp")
#
#export(TARGETS ${TARGET_NAME} NAMESPACE openvino::
# APPEND FILE "${CMAKE_BINARY_DIR}/OpenVINOTargets.cmake")

View File

@ -0,0 +1,45 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <openvino/core/variant.hpp>
#include <tensorflow_frontend/utility.hpp>
namespace ov {
namespace frontend {
class TF_API DecoderBase {
public:
/// \brief Get attribute value by name and requested type
///
/// \param name Attribute name
/// \param type_info Attribute type information
/// \return Shared pointer to appropriate value if it exists, 'nullptr' otherwise
virtual std::shared_ptr<ov::Variant> get_attribute(const std::string& name,
const VariantTypeInfo& type_info) const = 0;
/// \brief Get a number of inputs
virtual size_t get_input_size() const = 0;
/// \brief Get a producer name and its output port index
///
/// \param input_port_idx Input port index by which data is consumed
/// \param producer_name A producer name
/// \return producer_output_port_index Output port index from which data is generated
virtual void get_input_node(size_t input_port_idx,
std::string& producer_name,
size_t& producer_output_port_index) const = 0;
/// \brief Get operation type
virtual const std::string& get_op_type() const = 0;
/// \brief Get node name
virtual const std::string& get_op_name() const = 0;
/// \brief Destructor
virtual ~DecoderBase() = default;
};
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,35 @@
// Copyright (C) 2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <frontend_manager/frontend_exceptions.hpp>
#include <openvino/core/node.hpp>
namespace ov {
namespace frontend {
namespace tf {
class NodeContext;
class OpValidationFailureTF : public ngraph::frontend::OpValidationFailure {
public:
OpValidationFailureTF(const CheckLocInfo& check_loc_info, const NodeContext& node, const std::string& explanation)
: OpValidationFailure(check_loc_info, get_error_msg_prefix_tf(node), explanation) {}
private:
static std::string get_error_msg_prefix_tf(const NodeContext& node);
};
} // namespace tf
} // namespace frontend
/// \brief Macro to check whether a boolean condition holds.
/// \param node_context Object of NodeContext class
/// \param cond Condition to check
/// \param ... Additional error message info to be added to the error message via the `<<`
/// stream-insertion operator. Note that the expressions here will be evaluated lazily,
/// i.e., only if the `cond` evalutes to `false`.
/// \throws ::ov::OpValidationFailureTF if `cond` is false.
#define TF_OP_VALIDATION_CHECK(node_context, ...) \
NGRAPH_CHECK_HELPER(::ov::frontend::tf::OpValidationFailureTF, (node_context), __VA_ARGS__)
} // namespace ov

View File

@ -0,0 +1,84 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <frontend_manager/frontend.hpp>
#include <frontend_manager/input_model.hpp>
#include <functional>
#include <map>
#include <openvino/core/node_vector.hpp>
#include <openvino/core/variant.hpp>
#include <tensorflow_frontend/model.hpp>
#include <tensorflow_frontend/utility.hpp>
namespace ov {
namespace frontend {
namespace tf {
class NodeContext;
}
} // namespace frontend
} // namespace ov
namespace ov {
namespace frontend {
class TF_API FrontEndTF : public ngraph::frontend::FrontEnd {
public:
using CreatorFunction = std::function<::ov::OutputVector(const ::ov::frontend::tf::NodeContext&)>;
using TranslatorDictionaryType = std::map<const std::string, const CreatorFunction>;
private:
TranslatorDictionaryType m_op_translators;
public:
FrontEndTF();
/// \brief Completely convert the model
/// \return fully converted nGraph function
std::shared_ptr<Function> convert(ngraph::frontend::InputModel::Ptr model) const override;
/// \brief Completely convert the remaining, not converted part of a function.
/// \param partiallyConverted partially converted nGraph function
void convert(std::shared_ptr<Function> partiallyConverted) const override;
/// \brief Convert only those parts of the model that can be converted leaving others
/// as-is. Converted parts are not normalized by additional transformations; normalize
/// function or another form of convert function should be called to finalize the
/// conversion process.
/// \param model Input model
/// \return partially converted nGraph function
std::shared_ptr<Function> convert_partially(ngraph::frontend::InputModel::Ptr model) const override;
/// \brief Convert operations with one-to-one mapping with decoding nodes.
/// Each decoding node is an nGraph node representing a single FW operation node with
/// all attributes represented in FW-independent way.
/// \param model Input model
/// \return nGraph function after decoding
std::shared_ptr<Function> decode(ngraph::frontend::InputModel::Ptr model) const override;
/// \brief Runs normalization passes on function that was loaded with partial conversion
/// \param function partially converted nGraph function
void normalize(std::shared_ptr<ov::Function> function) const override;
/// \brief Gets name of this FrontEnd. Can be used by clients
std::string get_name() const override {
return "tf";
}
protected:
/// \brief Check if FrontEndTensorflow can recognize model from given parts
bool supported_impl(const std::vector<std::shared_ptr<ov::Variant>>& variants) const override;
ngraph::frontend::InputModel::Ptr load_impl(
const std::vector<std::shared_ptr<ov::Variant>>& variants) const override;
private:
void translate_graph(const std::shared_ptr<InputModelTF>& model,
const std::string& model_name,
bool fail_fast,
bool no_conversion,
std::shared_ptr<ov::Function>& ng_function) const;
};
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,36 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <tensorflow_frontend/decoder.hpp>
#include <tensorflow_frontend/utility.hpp>
namespace ov {
namespace frontend {
/// Abstract representation for an input model graph that gives nodes in topologically sorted order
class TF_API GraphIterator {
public:
using Ptr = std::shared_ptr<GraphIterator>;
/// \brief Get a number of operation nodes in the graph
virtual size_t size() const = 0;
/// \brief Set iterator to the start position
virtual void reset() = 0;
/// \brief Move to the next node in the graph
virtual void next() = 0;
/// \brief Returns true if iterator goes out of the range of available nodes
virtual bool is_end() const = 0;
/// \brief Return a pointer to a decoder of the current node
virtual std::shared_ptr<DecoderBase> get_decoder() const = 0;
/// \brief Destructor
virtual ~GraphIterator() = default;
};
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,43 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <frontend_manager/input_model.hpp>
#include <frontend_manager/place.hpp>
#include <tensorflow_frontend/graph_iterator.hpp>
#include <tensorflow_frontend/utility.hpp>
namespace ov {
namespace frontend {
class OpPlaceTF;
class TensorPlaceTF;
class TF_API InputModelTF : public ngraph::frontend::InputModel {
friend class FrontEndTF;
class InputModelTFImpl;
std::shared_ptr<InputModelTFImpl> _impl;
std::vector<std::shared_ptr<OpPlaceTF>> get_op_places() const;
std::map<std::string, std::shared_ptr<TensorPlaceTF>> get_tensor_places() const;
std::map<std::string, Output<Node>> get_tensor_values() const;
public:
explicit InputModelTF(const GraphIterator::Ptr& graph_iterator);
std::vector<ngraph::frontend::Place::Ptr> get_inputs() const override;
std::vector<ngraph::frontend::Place::Ptr> get_outputs() const override;
ngraph::frontend::Place::Ptr get_place_by_tensor_name(const std::string& tensorName) const override;
void override_all_outputs(const std::vector<ngraph::frontend::Place::Ptr>& outputs) override;
void override_all_inputs(const std::vector<ngraph::frontend::Place::Ptr>& inputs) override;
void extract_subgraph(const std::vector<ngraph::frontend::Place::Ptr>& inputs,
const std::vector<ngraph::frontend::Place::Ptr>& outputs) override;
void set_partial_shape(ngraph::frontend::Place::Ptr place, const ov::PartialShape&) override;
ov::PartialShape get_partial_shape(ngraph::frontend::Place::Ptr place) const override;
void set_element_type(ngraph::frontend::Place::Ptr place, const ov::element::Type&) override;
void set_tensor_value(ngraph::frontend::Place::Ptr place, const void* value) override;
};
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,176 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <frontend_manager/frontend.hpp>
#include <tensorflow_frontend/decoder.hpp>
namespace ov {
namespace frontend {
class TensorPlaceTF;
class OpPlaceTF;
class PlaceTF : public ngraph::frontend::Place {
public:
PlaceTF(const ngraph::frontend::InputModel& input_model, const std::vector<std::string>& names)
: m_input_model(input_model),
m_names(names) {}
explicit PlaceTF(const ngraph::frontend::InputModel& input_model)
: PlaceTF(input_model, std::vector<std::string>{}) {}
~PlaceTF() override = default;
bool is_input() const override;
bool is_output() const override;
bool is_equal(Ptr another) const override {
return this == another.get();
}
std::vector<std::string> get_names() const override {
return m_names;
}
private:
const ngraph::frontend::InputModel& m_input_model;
std::vector<std::string> m_names;
};
class InPortPlaceTF : public PlaceTF {
public:
explicit InPortPlaceTF(const ngraph::frontend::InputModel& input_model) : PlaceTF(input_model) {}
void set_op(const std::weak_ptr<OpPlaceTF>& op) {
m_op = op;
}
void set_source_tensor(const std::weak_ptr<TensorPlaceTF>& source_tensor);
// Internal usage
std::shared_ptr<TensorPlaceTF> get_source_tensor_tf() const;
std::shared_ptr<OpPlaceTF> get_op();
// External usage
std::vector<Ptr> get_consuming_operations() const override;
Ptr get_producing_operation() const override;
ngraph::frontend::Place::Ptr get_source_tensor() const override;
Ptr get_producing_port() const override;
bool is_equal_data(Ptr another) const override;
private:
std::weak_ptr<TensorPlaceTF> m_source_tensor;
std::weak_ptr<OpPlaceTF> m_op;
};
class OutPortPlaceTF : public PlaceTF {
public:
explicit OutPortPlaceTF(const ngraph::frontend::InputModel& input_model) : PlaceTF(input_model) {}
void set_op(const std::weak_ptr<OpPlaceTF>& op) {
m_op = op;
}
void set_target_tensor(const std::weak_ptr<TensorPlaceTF>& target_tensor);
std::shared_ptr<TensorPlaceTF> get_target_tensor_tf() const;
// External usage
std::vector<Ptr> get_consuming_operations() const override;
ngraph::frontend::Place::Ptr get_producing_operation() const override;
std::vector<ngraph::frontend::Place::Ptr> get_consuming_ports() const override;
Ptr get_target_tensor() const override;
bool is_equal_data(Ptr another) const override;
private:
std::weak_ptr<OpPlaceTF> m_op;
std::weak_ptr<TensorPlaceTF> m_target_tensor;
};
class OpPlaceTF : public PlaceTF {
public:
OpPlaceTF(const ngraph::frontend::InputModel& input_model, std::shared_ptr<DecoderBase> op_decoder);
void add_in_port(const std::shared_ptr<InPortPlaceTF>& input, const std::string& name);
void add_out_port(const std::shared_ptr<OutPortPlaceTF>& output, int idx);
// Internal usage
const std::vector<std::shared_ptr<OutPortPlaceTF>>& get_output_ports() const;
const std::map<std::string, std::vector<std::shared_ptr<InPortPlaceTF>>>& get_input_ports() const;
std::shared_ptr<InPortPlaceTF> get_input_port_tf(const std::string& inputName, int inputPortIndex) const;
std::shared_ptr<DecoderBase> get_decoder() const;
// External API methods
std::vector<ngraph::frontend::Place::Ptr> get_consuming_ports() const override;
Ptr get_output_port() const override;
Ptr get_output_port(int outputPortIndex) const override;
Ptr get_input_port() const override;
Ptr get_input_port(int inputPortIndex) const override;
Ptr get_input_port(const std::string& inputName) const override;
Ptr get_input_port(const std::string& inputName, int inputPortIndex) const override;
std::vector<Ptr> get_consuming_operations() const override;
std::vector<Ptr> get_consuming_operations(int outputPortIndex) const override;
Ptr get_producing_operation() const override;
Ptr get_producing_operation(int inputPortIndex) const override;
Ptr get_producing_operation(const std::string& inputName) const override;
Ptr get_producing_operation(const std::string& inputName, int inputPortIndex) const override;
Ptr get_source_tensor() const override;
Ptr get_source_tensor(int inputPortIndex) const override;
Ptr get_source_tensor(const std::string& inputName) const override;
Ptr get_source_tensor(const std::string& inputName, int inputPortIndex) const override;
Ptr get_target_tensor() const override;
Ptr get_target_tensor(int outputPortIndex) const override;
private:
std::shared_ptr<DecoderBase> m_op_decoder;
std::map<std::string, std::vector<std::shared_ptr<InPortPlaceTF>>> m_input_ports;
std::vector<std::shared_ptr<OutPortPlaceTF>> m_output_ports;
};
class TensorPlaceTF : public PlaceTF {
public:
TensorPlaceTF(const ngraph::frontend::InputModel& input_model,
const ov::PartialShape& pshape,
ov::element::Type type,
const std::vector<std::string>& names);
void add_producing_port(const std::shared_ptr<OutPortPlaceTF>& out_port);
void add_consuming_port(const std::shared_ptr<InPortPlaceTF>& in_port);
// Internal usage
const PartialShape& get_partial_shape() const {
return m_pshape;
}
const element::Type& get_element_type() const {
return m_type;
}
void set_partial_shape(const PartialShape& pshape) {
m_pshape = pshape;
}
void set_element_type(const element::Type& type) {
m_type = type;
}
// External usage
Ptr get_producing_operation() const override;
std::vector<ngraph::frontend::Place::Ptr> get_consuming_operations() const override;
std::vector<ngraph::frontend::Place::Ptr> get_consuming_ports() const override;
Ptr get_producing_port() const override;
bool is_equal_data(Ptr another) const override;
private:
PartialShape m_pshape;
element::Type m_type;
std::vector<std::weak_ptr<OutPortPlaceTF>> m_producing_ports;
std::vector<std::weak_ptr<InPortPlaceTF>> m_consuming_ports;
};
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,15 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <frontend_manager/frontend_exceptions.hpp>
#ifdef tensorflow_ngraph_frontend_EXPORTS
# define TF_API NGRAPH_HELPER_DLL_EXPORT
#else
# define TF_API NGRAPH_HELPER_DLL_IMPORT
#endif // tensorflow_ngraph_frontend_EXPORTS
#define NGRAPH_VLOG(I) std::ostringstream()

View File

@ -0,0 +1,127 @@
// Copyright (C) 2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include "decoder_proto.hpp"
#include "node_context.hpp"
namespace ov {
namespace frontend {
namespace tf {
const std::map<::tensorflow::DataType, ov::element::Type>& TYPE_MAP() {
static const std::map<::tensorflow::DataType, ov::element::Type> type_map{
{::tensorflow::DataType::DT_BOOL, ov::element::boolean},
{::tensorflow::DataType::DT_INT16, ov::element::i16},
{::tensorflow::DataType::DT_INT32, ov::element::i32},
{::tensorflow::DataType::DT_INT64, ov::element::i64},
{::tensorflow::DataType::DT_HALF, ov::element::f16},
{::tensorflow::DataType::DT_FLOAT, ov::element::f32},
{::tensorflow::DataType::DT_DOUBLE, ov::element::f64},
{::tensorflow::DataType::DT_UINT8, ov::element::u8},
{::tensorflow::DataType::DT_INT8, ov::element::i8},
{::tensorflow::DataType::DT_BFLOAT16, ov::element::bf16}};
return type_map;
}
std::shared_ptr<ov::Variant> DecoderTFProto::get_attribute(const std::string& name,
const VariantTypeInfo& type_info) const {
auto attrs = decode_attribute_helper(name);
if (attrs.empty()) {
return nullptr;
}
if (type_info == VariantWrapper<std::string>::get_type_info_static()) {
return std::make_shared<VariantWrapper<std::string>>(attrs[0].s());
} else if (type_info == VariantWrapper<int64_t>::get_type_info_static()) {
return std::make_shared<VariantWrapper<int64_t>>(attrs[0].i());
} else if (type_info == VariantWrapper<std::vector<int64_t>>::get_type_info_static()) {
std::vector<int64_t> longs;
longs.reserve(attrs[0].list().i_size());
for (size_t idx = 0; idx < attrs[0].list().i_size(); ++idx) {
longs.push_back(attrs[0].list().i(idx));
}
return std::make_shared<VariantWrapper<std::vector<int64_t>>>(longs);
} else if (type_info == VariantWrapper<int32_t>::get_type_info_static()) {
return std::make_shared<VariantWrapper<int32_t>>(static_cast<int32_t>(attrs[0].i()));
} else if (type_info == VariantWrapper<std::vector<int32_t>>::get_type_info_static()) {
std::vector<int32_t> ints;
ints.reserve(attrs[0].list().i_size());
for (size_t idx = 0; idx < attrs[0].list().i_size(); ++idx) {
ints.push_back(static_cast<int32_t>(attrs[0].list().i(idx)));
}
return std::make_shared<VariantWrapper<std::vector<int32_t>>>(ints);
} else if (type_info == VariantWrapper<float>::get_type_info_static()) {
return std::make_shared<VariantWrapper<float>>(attrs[0].f());
} else if (type_info == VariantWrapper<std::vector<float>>::get_type_info_static()) {
std::vector<float> floats;
floats.reserve(attrs[0].list().i_size());
for (size_t idx = 0; idx < attrs[0].list().i_size(); ++idx) {
floats.push_back(attrs[0].list().f(idx));
}
return std::make_shared<VariantWrapper<std::vector<float>>>(floats);
} else if (type_info == VariantWrapper<ov::element::Type>::get_type_info_static()) {
auto data_type = attrs[0].type();
return std::make_shared<VariantWrapper<ov::element::Type>>(TYPE_MAP().at(data_type));
} else if (type_info == VariantWrapper<bool>::get_type_info_static()) {
return std::make_shared<VariantWrapper<bool>>(attrs[0].b());
} else if (type_info == VariantWrapper<::tensorflow::DataType>::get_type_info_static()) {
return std::make_shared<VariantWrapper<::tensorflow::DataType>>(attrs[0].type());
} else if (type_info == VariantWrapper<::tensorflow::TensorProto>::get_type_info_static()) {
return std::make_shared<VariantWrapper<::tensorflow::TensorProto>>(attrs[0].tensor());
} else if (type_info == VariantWrapper<::ov::PartialShape>::get_type_info_static()) {
std::vector<ov::Dimension> dims;
auto tf_shape = attrs[0].shape();
for (int i = 0; i < tf_shape.dim_size(); i++) {
dims.push_back(tf_shape.dim(i).size());
}
auto pshape = ov::PartialShape(dims);
return std::make_shared<VariantWrapper<::ov::PartialShape>>(pshape);
}
// type is not supported by decoder
return nullptr;
}
size_t DecoderTFProto::get_input_size() const {
return m_node_def->input_size();
}
void DecoderTFProto::get_input_node(size_t input_port_idx,
std::string& producer_name,
size_t& producer_output_port_index) const {
// TODO: handle body graph nodes with a couple of columns
std::string producer_port_name = m_node_def->input(input_port_idx);
auto delim_pos = producer_port_name.find(':');
if (delim_pos != std::string::npos) {
producer_name = producer_port_name.substr(0, delim_pos);
producer_output_port_index = std::stoi(producer_port_name.substr(delim_pos));
return;
}
producer_name = producer_port_name;
producer_output_port_index = 0;
}
const std::string& DecoderTFProto::get_op_type() const {
return m_node_def->op();
}
const std::string& DecoderTFProto::get_op_name() const {
return m_node_def->name();
}
std::vector<::tensorflow::AttrValue> DecoderTFProto::decode_attribute_helper(const std::string& name) const {
auto attr_map = m_node_def->attr();
FRONT_END_GENERAL_CHECK(attr_map.contains(name),
"An error occurred while parsing the ",
name,
" attribute of ",
this->get_op_type(),
"node");
auto value = m_node_def->attr().at(name);
return {value};
}
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,45 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <ngraph/ngraph.hpp>
#include <string>
#include <tensorflow_frontend/decoder.hpp>
#include <tensorflow_frontend/frontend.hpp>
#include <tensorflow_frontend/place.hpp>
#include <vector>
#include "attr_value.pb.h"
#include "node_def.pb.h"
#include "types.pb.h"
namespace ov {
namespace frontend {
namespace tf {
class DecoderTFProto : public DecoderBase {
public:
explicit DecoderTFProto(const ::tensorflow::NodeDef* node_def) : m_node_def(node_def) {}
std::shared_ptr<ov::Variant> get_attribute(const std::string& name,
const VariantTypeInfo& type_info) const override;
size_t get_input_size() const override;
void get_input_node(size_t input_port_idx,
std::string& producer_name,
size_t& producer_output_port_index) const override;
const std::string& get_op_type() const override;
const std::string& get_op_name() const override;
private:
std::vector<::tensorflow::AttrValue> decode_attribute_helper(const std::string& name) const;
const ::tensorflow::NodeDef* m_node_def;
};
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,19 @@
// Copyright (C) 2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <tensorflow_frontend/exceptions.hpp>
#include "node_context.hpp"
namespace ov {
namespace frontend {
namespace tf {
std::string OpValidationFailureTF::get_error_msg_prefix_tf(const tf::NodeContext& node) {
std::stringstream ss;
ss << "While validating node '" << node.get_op_type() << '\'';
return ss.str();
}
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,350 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <openvino/util/common_util.hpp>
#include <tensorflow_frontend/frontend.hpp>
#include <tensorflow_frontend/model.hpp>
#include "op_table.hpp"
#include "tf_framework_node.hpp"
#include "utils.hpp"
using namespace ::ov::frontend;
using namespace ::ov::frontend::tf;
namespace {
void translate_framework_node(const std::shared_ptr<TFFrameworkNode>& node,
const FrontEndTF::TranslatorDictionaryType& op_translators) {
auto type = node->get_op_type();
const auto& TRANSLATE_OP_MAP = op_translators;
auto translator_it = TRANSLATE_OP_MAP.find(type);
FRONT_END_OP_CONVERSION_CHECK(translator_it != TRANSLATE_OP_MAP.end(), "No translator found for ", type, " node.");
ov::OutputVector ng_inputs;
NamedInputs named_inputs;
size_t input_port_idx = 0;
for (const auto& input : node->input_values()) {
ng_inputs.push_back(input);
named_inputs[input_port_idx++] = {input};
}
NodeContext node_ctx(*node->get_decoder(), named_inputs);
auto new_node_outputs = translator_it->second(node_ctx);
SetTracingInfo(node_ctx.get_name(), new_node_outputs.front());
auto new_output = new_node_outputs.begin();
auto old_outputs = node->outputs();
auto old_output = old_outputs.begin();
for (; new_output != new_node_outputs.end() && old_output != old_outputs.end(); ++old_output, ++new_output) {
old_output->replace(*new_output);
}
}
} // namespace
FrontEndTF::FrontEndTF() : m_op_translators(tf::op::get_supported_ops()) {}
void FrontEndTF::translate_graph(const std::shared_ptr<InputModelTF>& model,
const std::string& model_name,
bool fail_fast,
bool no_conversion,
std::shared_ptr<ov::Function>& ng_function) const {
// a map from operation names to generated nGraph Output<TFNodeDecoder>
tf::OpMap ng_op_map;
ov::ParameterVector params;
ov::ResultVector results;
const auto& operation_places = model->get_op_places();
const auto& model_inputs = model->get_inputs();
const auto& model_outputs = model->get_outputs();
const auto& model_frozen_inputs = model->get_tensor_values();
std::map<const std::string, const std::function<ov::OutputVector(const NodeContext&)>> translate_map;
const auto& TRANSLATE_OP_MAP = m_op_translators;
if (no_conversion) {
const std::set<std::string> required_types{"Placeholder", "_Retval", "NoOp"};
for (const auto& name : required_types) {
translate_map.emplace(name, TRANSLATE_OP_MAP.at(name));
}
} else {
translate_map.insert(TRANSLATE_OP_MAP.begin(), TRANSLATE_OP_MAP.end());
}
// fill ng_op_map with Constant outputs for frozen inputs
for (const auto& frozen_input : model_frozen_inputs) {
const auto& frozen_input_name = frozen_input.first;
const auto& frozen_input_value = frozen_input.second;
FRONT_END_GENERAL_CHECK(ng_op_map.count(frozen_input_name) == 0,
"Input with frozen value has been already met: " + frozen_input_name);
ng_op_map[frozen_input_name] = {frozen_input_value};
}
// create parameter nodes for all tensor places corresponding to inputs
for (const auto& input_place : model_inputs) {
FRONT_END_GENERAL_CHECK(input_place->get_names().size() == 1, "Input place must have one name.");
auto input_name = input_place->get_names()[0];
if (ng_op_map.count(input_name)) {
// probably this input is frozen
continue;
}
const auto& input_tensor_place = std::dynamic_pointer_cast<TensorPlaceTF>(input_place);
auto input_shape = input_tensor_place->get_partial_shape();
auto input_type = input_tensor_place->get_element_type();
auto input_ng_output = ConstructNgNode<ov::opset8::Parameter>(input_name, input_type, input_shape);
auto input_ng_node = std::dynamic_pointer_cast<ov::opset8::Parameter>(input_ng_output.get_node_shared_ptr());
params.push_back(input_ng_node);
ng_op_map[input_name] = {input_ng_output};
}
// create the nGraph ops from TensorFlow ops
for (const auto& operation_place : operation_places) {
auto operation_decoder = operation_place->get_decoder();
auto operation_name = operation_place->get_names()[0];
// output for parameter nodes has been already generated
if (ng_op_map.count(operation_name)) {
continue;
}
// prepare a list of nGraph node inputs for each node
ov::OutputVector ng_inputs;
::ov::frontend::tf::NamedInputs named_inputs;
for (size_t input_port_idx = 0; input_port_idx < operation_decoder->get_input_size(); ++input_port_idx) {
std::string producer_name;
size_t producer_port_idx;
try {
operation_decoder->get_input_node(input_port_idx, producer_name, producer_port_idx);
} catch (const std::exception& e) {
FRONT_END_THROW("[ ERROR ] Exception happened when preparing input " + std::to_string(input_port_idx) +
" for op '" + operation_decoder->get_op_name() + "', expected input name: '" +
producer_name + "', expected input port index: " + std::to_string(producer_port_idx) +
'\n');
}
// TODO: re-implement the logic below once Place graph structure is implemented
// Using Place graph structure (OpPlace, In/OutPortPlace places and their connections) can give
// names of ports and operations that can be used for further check about existence in ng_op_map
// check if output vector for places have been already defined and the order of this check is important
// it moves from places corresponding to input port of the current operation node to output port of original
// producers
if (ng_op_map.count(std::to_string(input_port_idx) + ":" + operation_name)) {
const auto& input_outputs_vector = ng_op_map.at(std::to_string(input_port_idx) + ":" + operation_name);
FRONT_END_GENERAL_CHECK(input_outputs_vector.size() == 1,
"Input created with pruning must have one output");
ng_inputs.push_back(input_outputs_vector.at(0));
named_inputs[input_port_idx] = {input_outputs_vector.at(0)};
} else if (ng_op_map.count(producer_name + ":" + std::to_string(producer_port_idx))) {
const auto& input_outputs_vector =
ng_op_map.at(producer_name + ":" + std::to_string(producer_port_idx));
FRONT_END_GENERAL_CHECK(input_outputs_vector.size() == 1,
"Input created with pruning must have one output");
ng_inputs.push_back(input_outputs_vector.at(0));
named_inputs[input_port_idx] = {input_outputs_vector.at(0)};
} else if (ng_op_map.count(producer_name)) {
const auto& input_outputs_vector = ng_op_map.at(producer_name);
FRONT_END_GENERAL_CHECK(input_outputs_vector.size() > producer_port_idx,
"Input created with pruning must have one output");
ng_inputs.push_back(input_outputs_vector.at(producer_port_idx));
named_inputs[input_port_idx] = {input_outputs_vector.at(producer_port_idx)};
} else {
FRONT_END_GENERAL_CHECK(false,
"No input is found for node \"" + operation_name + "\" by port" +
std::to_string(producer_port_idx));
}
}
// generate nGraph node output vector for the current operation node
ov::OutputVector ng_outputs;
try {
FRONT_END_OP_CONVERSION_CHECK(translate_map.count(operation_decoder->get_op_type()),
"No translator found for " + operation_decoder->get_op_type() + " node.");
auto op_fun = &(translate_map[operation_decoder->get_op_type()]);
// NodeContext node_context(ng_inputs, operation_decoder, model_inputs);
// TODO: Check why NodeContextNew doesn't have ngOutputVector ng_inputs input in constructor
::ov::frontend::tf::NodeContext node_context(*operation_decoder, named_inputs);
// generate nGraph node output vector using translator for given operation type
ng_outputs = (*op_fun)(node_context);
} catch (...) {
if (fail_fast) {
// re-throw any exception
throw;
} else {
auto ng_node = std::make_shared<TFFrameworkNode>(operation_decoder,
ng_inputs,
operation_place->get_output_ports().size());
SetTracingInfo(operation_name, ng_node);
ng_outputs = ng_node->outputs();
}
}
// register nGraph node outputs in the map for new operation node
for (const auto& output : ng_outputs) {
if (auto result = std::dynamic_pointer_cast<ov::opset8::Result>(output.get_node_shared_ptr())) {
// do not add RetVal type operation to ng_op_map
results.push_back(result);
} else {
auto param = std::dynamic_pointer_cast<ov::opset8::Parameter>(output.get_node_shared_ptr());
if (param && operation_decoder->get_op_type() != "Identity") {
params.push_back(param);
}
ng_op_map[operation_name].push_back(output);
}
}
}
// create Result nodes for all model outputs
for (const auto& model_output : model_outputs) {
auto model_output_tensor_place = std::dynamic_pointer_cast<TensorPlaceTF>(model_output);
auto model_output_name = model_output_tensor_place->get_names()[0];
std::string operation_name;
std::string port_type;
size_t port_index;
ov::frontend::tf::extract_operation_name_and_port(model_output_name, operation_name, port_index, port_type);
if (port_type == "none") {
for (const auto& node_output : ng_op_map[operation_name]) {
results.push_back(std::make_shared<ov::opset8::Result>(node_output));
}
} else if (port_type == "out") {
const auto& node_outputs = ng_op_map[operation_name];
FRONT_END_GENERAL_CHECK(node_outputs.size() > port_index,
"Output port with index " + std::to_string(port_index) + " of " + operation_name +
"node specified as custom output does not exist");
results.push_back(std::make_shared<ov::opset8::Result>(node_outputs[port_index]));
} else if (port_type == "in") {
// TODO: avoid this traversing by having a map for OpPlace objects, for example
std::shared_ptr<OpPlaceTF> operation_place = nullptr;
for (const auto& op_place : operation_places) {
FRONT_END_GENERAL_CHECK(!op_place->get_names().empty(), "No names for OpPlace found.");
if (op_place->get_names()[0] == operation_name) {
operation_place = op_place;
}
}
FRONT_END_GENERAL_CHECK(operation_place, "There is no operation place with a name: " + operation_name);
auto operation_decoder = operation_place->get_decoder();
// get to know a producer node and by which its output port data is generated
std::string producer_name;
size_t producer_port_idx;
try {
operation_decoder->get_input_node(port_index, producer_name, producer_port_idx);
} catch (const std::exception& e) {
FRONT_END_THROW("[ ERROR ] Exception happened when preparing input " + std::to_string(port_index) +
" for op '" + operation_decoder->get_op_name() + "', expected input name: '" +
producer_name + "', expected input port index: " + std::to_string(producer_port_idx) +
'\n');
}
// add Result node for this producer output port
const auto& node_outputs = ng_op_map[producer_name];
FRONT_END_GENERAL_CHECK(node_outputs.size() > producer_port_idx,
"Output port with index " + std::to_string(producer_port_idx) + " of " +
producer_name + "node specified as custom output does not exist");
results.push_back(std::make_shared<ov::opset8::Result>(node_outputs[producer_port_idx]));
}
}
// find all terminal nodes in ngraph graph to complete list of results
if (results.empty()) {
for (const auto& node_output_vector : ng_op_map) {
for (const auto& output : node_output_vector.second) {
if (output.get_target_inputs().empty() &&
!std::dynamic_pointer_cast<ov::opset8::Result>(output.get_node_shared_ptr())) {
results.push_back(std::make_shared<ov::opset8::Result>(output));
}
}
}
}
// TODO: reorder results and params according to indices given in RT info (if any)
// create the nGraph function
ng_function = std::make_shared<ov::Function>(results, params, model_name);
NGRAPH_VLOG(5) << "Done with translations";
}
/// \brief Check if FrontEndTensorflow can recognize model from given parts
bool FrontEndTF::supported_impl(const std::vector<std::shared_ptr<ov::Variant>>& variants) const {
// TODO: Support other TensorFlow formats: SavedModel, .meta, checkpoint, pbtxt
if (variants.size() != 1)
return false;
// Validating first path, it must contain a model
if (ov::is_type<VariantWrapper<std::string>>(variants[0])) {
std::string suffix = ".pb";
std::string model_path = ov::as_type_ptr<VariantWrapper<std::string>>(variants[0])->get();
if (ov::util::ends_with(model_path, suffix.c_str())) {
return true;
}
}
return false;
}
ngraph::frontend::InputModel::Ptr FrontEndTF::load_impl(
const std::vector<std::shared_ptr<ov::Variant>>& variants) const {
// TODO: Support other TensorFlow formats: SavedModel, .meta, checkpoint, pbtxt
if (variants.size() == 1) {
// a case when binary protobuf format is provided
if (ov::is_type<VariantWrapper<std::string>>(variants[0])) {
std::string suffix = ".pb";
std::string model_path = ov::as_type_ptr<VariantWrapper<std::string>>(variants[0])->get();
if (ov::util::ends_with(model_path, suffix.c_str())) {
return std::make_shared<InputModelTF>(
std::make_shared<::ov::frontend::tf::GraphIteratorProto>(model_path));
}
}
}
return nullptr;
}
std::shared_ptr<ov::Function> FrontEndTF::convert(ngraph::frontend::InputModel::Ptr model) const {
auto model_tf = std::dynamic_pointer_cast<InputModelTF>(model);
std::shared_ptr<ov::Function> f;
translate_graph(model_tf, "here_should_be_a_graph_name", true, false, f);
normalize(f);
// TODO: check that nGraph function does not contain operations which are not in the opset
return f;
}
std::shared_ptr<ov::Function> FrontEndTF::convert_partially(ngraph::frontend::InputModel::Ptr model) const {
auto model_tf = std::dynamic_pointer_cast<InputModelTF>(model);
std::shared_ptr<ov::Function> f;
translate_graph(model_tf, "here_should_be_a_graph_name", false, false, f);
normalize(f);
return f;
}
std::shared_ptr<ov::Function> FrontEndTF::decode(ngraph::frontend::InputModel::Ptr model) const {
auto model_tf = std::dynamic_pointer_cast<InputModelTF>(model);
std::shared_ptr<ov::Function> f;
translate_graph(model_tf, "here_should_be_a_graph_name", false, true, f);
return f;
}
void FrontEndTF::convert(std::shared_ptr<ov::Function> partiallyConverted) const {
for (const auto& node : partiallyConverted->get_ordered_ops()) {
if (ov::is_type<TFFrameworkNode>(node)) {
translate_framework_node(std::dynamic_pointer_cast<TFFrameworkNode>(node), m_op_translators);
}
}
for (const auto& result : partiallyConverted->get_results()) {
result->validate_and_infer_types();
}
normalize(partiallyConverted);
}
void FrontEndTF::normalize(std::shared_ptr<ov::Function> function) const {
ov::pass::Manager manager;
// TODO: switch on TransposeSinking once it is ready
// manager.register_pass<ov::pass::TransposeSinking>();
manager.run_passes(function);
}

View File

@ -0,0 +1,61 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <fstream>
#include <tensorflow_frontend/decoder.hpp>
#include <tensorflow_frontend/graph_iterator.hpp>
#include "decoder_proto.hpp"
#include "graph.pb.h"
#include "node_def.pb.h"
namespace ov {
namespace frontend {
namespace tf {
class GraphIteratorProto : public GraphIterator {
std::vector<const ::tensorflow::NodeDef*> m_nodes;
size_t node_index = 0;
std::shared_ptr<::tensorflow::GraphDef> m_graph_def;
public:
template <typename T>
GraphIteratorProto(const std::basic_string<T>& path) : m_graph_def(std::make_shared<::tensorflow::GraphDef>()) {
std::ifstream pb_stream(path, std::ios::in | std::ifstream::binary);
FRONT_END_GENERAL_CHECK(pb_stream && pb_stream.is_open(), "Model file does not exist");
FRONT_END_GENERAL_CHECK(m_graph_def->ParseFromIstream(&pb_stream), "Model cannot be parsed");
m_nodes.resize(m_graph_def->node_size());
for (size_t i = 0; i < m_nodes.size(); ++i)
m_nodes[i] = &m_graph_def->node(i);
}
/// Set iterator to the start position
void reset() override {
node_index = 0;
}
size_t size() const override {
return m_nodes.size();
}
/// Moves to the next node in the graph
void next() override {
node_index++;
}
bool is_end() const override {
return node_index >= m_nodes.size();
}
/// Return NodeContext for the current node that iterator points to
std::shared_ptr<DecoderBase> get_decoder() const override {
return std::make_shared<DecoderTFProto>(m_nodes[node_index]);
}
};
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,387 @@
// Copyright (C) 2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <frontend_manager/frontend_exceptions.hpp>
#include <fstream>
#include <openvino/opsets/opset7.hpp>
#include <queue>
#include <tensorflow_frontend/graph_iterator.hpp>
#include <tensorflow_frontend/model.hpp>
#include <tensorflow_frontend/place.hpp>
#include <tensorflow_frontend/utility.hpp>
#include "graph_iterator_proto.hpp"
#include "ngraph_conversions.hpp"
#include "node_context.hpp"
#include "utils.hpp"
using namespace google;
namespace ov {
namespace frontend {
namespace tf {
void extract_operation_name_and_port(const std::string& port_name,
std::string& operation_name,
size_t& port_index,
std::string& port_type) {
constexpr char delimeter[] = ":";
auto pos = port_name.find(delimeter);
if (pos == std::string::npos) {
operation_name = port_name;
port_type = "none";
port_index = 0;
return;
}
FRONT_END_GENERAL_CHECK((0 < pos) && (pos + 1 < port_name.length()), "Incorrect port name specified: " + port_name);
auto left_part = port_name.substr(0, pos);
auto right_part = port_name.substr(pos + 1, port_name.length() - pos);
if (left_part.find_first_not_of("0123456789") == std::string::npos) {
port_type = "in";
operation_name = right_part;
port_index = std::atoi(left_part.c_str());
} else if (right_part.find_first_not_of("0123456789") == std::string::npos) {
port_type = "out";
operation_name = left_part;
port_index = std::atoi(right_part.c_str());
} else {
FRONT_END_GENERAL_CHECK(false, "Incorrect port name specified: " + port_name);
}
}
} // namespace tf
class InputModelTF::InputModelTFImpl {
public:
InputModelTFImpl(const GraphIterator::Ptr& graph_iterator, const ngraph::frontend::InputModel& input_model);
std::vector<ngraph::frontend::Place::Ptr> getInputs() const;
std::vector<ngraph::frontend::Place::Ptr> getOutputs() const;
ngraph::frontend::Place::Ptr getPlaceByTensorName(const std::string& tensorName) const;
void overrideAllOutputs(const std::vector<ngraph::frontend::Place::Ptr>& outputs);
void overrideAllInputs(const std::vector<ngraph::frontend::Place::Ptr>& inputs);
void extractSubgraph(const std::vector<ngraph::frontend::Place::Ptr>& inputs,
const std::vector<ngraph::frontend::Place::Ptr>& outputs);
void setPartialShape(ngraph::frontend::Place::Ptr place, const ov::PartialShape&);
ov::PartialShape getPartialShape(ngraph::frontend::Place::Ptr place) const;
void setElementType(ngraph::frontend::Place::Ptr place, const ov::element::Type&);
void setTensorValue(ngraph::frontend::Place::Ptr place, const void* value);
std::vector<std::shared_ptr<OpPlaceTF>> get_op_places() const;
std::map<std::string, std::shared_ptr<TensorPlaceTF>> get_tensor_places() const {
return m_tensor_places;
}
std::map<std::string, Output<Node>> get_tensor_values() const {
return m_tensor_values;
};
private:
void loadPlaces();
std::vector<std::shared_ptr<OpPlaceTF>> determine_cut_nodes() const;
std::vector<std::shared_ptr<OpPlaceTF>> m_op_places;
std::map<std::string, std::shared_ptr<OpPlaceTF>> m_op_places_map;
mutable std::map<std::string, std::shared_ptr<TensorPlaceTF>> m_tensor_places;
std::vector<ngraph::frontend::Place::Ptr> m_inputs;
std::vector<ngraph::frontend::Place::Ptr> m_outputs;
std::map<std::string, Output<Node>> m_tensor_values;
std::shared_ptr<GraphIterator> m_graph_iterator;
const ngraph::frontend::InputModel& m_input_model;
// shows if some nodes might be deleted from graph
bool m_graph_changed = false;
};
void InputModelTF::InputModelTFImpl::loadPlaces() {
std::set<std::string> all_op_names;
std::set<std::string> op_names_with_consumers;
m_inputs.clear();
for (; !m_graph_iterator->is_end(); m_graph_iterator->next()) {
auto node_decoder = m_graph_iterator->get_decoder();
auto op_name = node_decoder->get_op_name();
auto op_type = node_decoder->get_op_type();
auto op_place = std::make_shared<OpPlaceTF>(m_input_model, node_decoder);
all_op_names.insert(op_name);
m_op_places.push_back(op_place);
m_op_places_map[op_name] = op_place;
if (op_type == "Placeholder") {
auto pshape = std::dynamic_pointer_cast<VariantWrapper<ov::PartialShape>>(
node_decoder->get_attribute("shape", VariantWrapper<ov::PartialShape>::get_type_info_static()));
auto type = std::dynamic_pointer_cast<VariantWrapper<ov::element::Type>>(
node_decoder->get_attribute("dtype", VariantWrapper<ov::element::Type>::get_type_info_static()));
std::vector<std::string> names = {op_name};
auto tensor_place = std::make_shared<TensorPlaceTF>(m_input_model, pshape->get(), type->get(), names);
m_tensor_places[op_name] = tensor_place;
m_inputs.push_back(tensor_place);
}
for (size_t input_port_idx = 0; input_port_idx < node_decoder->get_input_size(); ++input_port_idx) {
std::string producer_op_name;
size_t producer_output_port_idx;
try {
node_decoder->get_input_node(input_port_idx, producer_op_name, producer_output_port_idx);
op_names_with_consumers.insert(producer_op_name);
} catch (const std::exception& e) {
FRONT_END_THROW("[ ERROR ] Exception happened when preparing input " + std::to_string(input_port_idx) +
" for op '" + node_decoder->get_op_name() + "', expected input name: '" +
producer_op_name +
"', expected input port index: " + std::to_string(producer_output_port_idx));
}
}
}
std::set<std::string> op_names_without_consumers;
std::set_difference(all_op_names.begin(),
all_op_names.end(),
op_names_with_consumers.begin(),
op_names_with_consumers.end(),
std::inserter(op_names_without_consumers, op_names_without_consumers.begin()));
m_graph_iterator->reset();
m_outputs.clear();
for (auto& output_name : op_names_without_consumers) {
std::vector<std::string> output_names = {output_name};
auto output_place =
std::make_shared<TensorPlaceTF>(m_input_model, ov::PartialShape({}), ov::element::undefined, output_names);
m_tensor_places[output_name] = output_place;
m_outputs.push_back(output_place);
}
}
std::vector<std::shared_ptr<OpPlaceTF>> InputModelTF::InputModelTFImpl::get_op_places() const {
if (m_graph_changed) {
return determine_cut_nodes();
}
return m_op_places;
}
std::vector<std::shared_ptr<OpPlaceTF>> InputModelTF::InputModelTFImpl::determine_cut_nodes() const {
std::queue<std::shared_ptr<DecoderBase>> decoders_queue;
std::unordered_set<std::string> visited;
std::vector<std::shared_ptr<OpPlaceTF>> new_ops;
for (const auto& output_place : m_outputs) {
FRONT_END_GENERAL_CHECK(output_place->get_names().size() > 0, "TensorPlace must have at least one name.");
auto output_place_name = output_place->get_names()[0];
std::string operation_name;
size_t port_idx;
std::string port_type;
tf::extract_operation_name_and_port(output_place_name, operation_name, port_idx, port_type);
if (!visited.count(operation_name)) {
visited.insert(operation_name);
FRONT_END_GENERAL_CHECK(m_op_places_map.count(operation_name),
"Custom specified output is incorrect: " + output_place_name);
auto output_operation_place = m_op_places_map.at(operation_name);
FRONT_END_GENERAL_CHECK(output_operation_place,
"There is not operation place in the map: " + operation_name);
new_ops.push_back(output_operation_place);
decoders_queue.push(output_operation_place->get_decoder());
}
}
while (!decoders_queue.empty()) {
auto operation_decoder = decoders_queue.front();
decoders_queue.pop();
auto current_operation_name = operation_decoder->get_op_name();
for (size_t input_port_idx = 0; input_port_idx < operation_decoder->get_input_size(); ++input_port_idx) {
std::string producer_name;
size_t producer_output_port_idx;
try {
operation_decoder->get_input_node(input_port_idx, producer_name, producer_output_port_idx);
} catch (const std::exception& e) {
FRONT_END_THROW("[ ERROR ] Exception happened when preparing input " + std::to_string(input_port_idx) +
" for op '" + operation_decoder->get_op_name() + "', expected input name: '" +
producer_name +
"', expected input port index: " + std::to_string(producer_output_port_idx) + '\n');
}
// TODO: re-implement the logic below using Place graph structure (with OpPlace, In/OutPortPlace
// connections) and based on check if Place->is_input() decide to leave a node or not
// is_input is a flag to leave producer operation node or not.
// this producing node is not left if consumer is pruned by its input port,
// the producer node is pruned by its output port or the producer becomes new input
// 1. check if the current node is pruned by its input port
bool is_input = false;
std::string input_port_name = std::to_string(input_port_idx) + ":" + current_operation_name;
if (m_tensor_places.find(input_port_name) != m_tensor_places.end()) {
const auto& tensor_place = m_tensor_places[input_port_name];
is_input = is_input || (tensor_place->is_input() ? true : false);
}
// 2. check if the producer node is pruned by its output port
std::string output_port_name = producer_name + ":" + std::to_string(producer_output_port_idx);
if (m_tensor_places.find(output_port_name) != m_tensor_places.end()) {
const auto& tensor_place = m_tensor_places[output_port_name];
is_input = is_input || (tensor_place->is_input() ? true : false);
}
// 3. check if the current node is an input
FRONT_END_GENERAL_CHECK(m_op_places_map.count(producer_name),
"There is no operation node with name: " + producer_name);
const auto& producer_operation_place = m_op_places_map.at(producer_name);
if (m_tensor_places.find(producer_name) != m_tensor_places.end()) {
const auto& tensor_place = m_tensor_places[producer_name];
is_input |= (tensor_place->is_input() ? true : false);
}
if (!is_input && !visited.count(producer_name)) {
visited.insert(producer_name);
new_ops.push_back(producer_operation_place);
decoders_queue.push(producer_operation_place->get_decoder());
}
}
}
std::reverse(new_ops.begin(), new_ops.end());
return new_ops;
}
InputModelTF::InputModelTFImpl::InputModelTFImpl(const GraphIterator::Ptr& graph_iterator,
const ngraph::frontend::InputModel& input_model)
: m_input_model(input_model),
m_graph_iterator(graph_iterator) {
FRONT_END_GENERAL_CHECK(m_graph_iterator, "Null pointer specified for GraphIterator");
loadPlaces();
}
std::vector<ngraph::frontend::Place::Ptr> InputModelTF::InputModelTFImpl::getInputs() const {
return m_inputs;
}
std::vector<ngraph::frontend::Place::Ptr> InputModelTF::InputModelTFImpl::getOutputs() const {
return m_outputs;
}
ngraph::frontend::Place::Ptr InputModelTF::InputModelTFImpl::getPlaceByTensorName(const std::string& tensorName) const {
if (m_tensor_places.find(tensorName) != m_tensor_places.end())
return m_tensor_places.at(tensorName);
// check that operation node exists for which this place is specified
std::string operation_name;
size_t port_idx;
std::string port_type;
tf::extract_operation_name_and_port(tensorName, operation_name, port_idx, port_type);
if (m_op_places_map.find(operation_name) != m_op_places_map.end()) {
std::vector<std::string> names = {tensorName};
auto m_var_place =
std::make_shared<TensorPlaceTF>(m_input_model, ov::PartialShape(), ov::element::undefined, names);
m_tensor_places[tensorName] = m_var_place;
return m_var_place;
}
return nullptr;
}
std::shared_ptr<TensorPlaceTF> castToTensorPlace(const ngraph::frontend::Place::Ptr& place) {
if (auto var_place = std::dynamic_pointer_cast<TensorPlaceTF>(place)) {
return var_place;
} else if (auto in_port_place = std::dynamic_pointer_cast<InPortPlaceTF>(place)) {
return in_port_place->get_source_tensor_tf();
} else if (auto out_port_place = std::dynamic_pointer_cast<OutPortPlaceTF>(place)) {
return out_port_place->get_target_tensor_tf();
}
FRONT_END_GENERAL_CHECK(false, "Cannot cast this Place to TensorPlaceTF.");
}
void InputModelTF::InputModelTFImpl::overrideAllInputs(const std::vector<ngraph::frontend::Place::Ptr>& inputs) {
m_graph_changed = true;
m_inputs.clear();
for (const auto& input_place : inputs) {
m_inputs.push_back(castToTensorPlace(input_place));
}
}
void InputModelTF::InputModelTFImpl::overrideAllOutputs(const std::vector<ngraph::frontend::Place::Ptr>& outputs) {
m_graph_changed = true;
m_outputs.clear();
for (const auto& output_place : outputs) {
m_outputs.push_back(castToTensorPlace(output_place));
}
}
void InputModelTF::InputModelTFImpl::extractSubgraph(const std::vector<ngraph::frontend::Place::Ptr>& inputs,
const std::vector<ngraph::frontend::Place::Ptr>& outputs) {
m_graph_changed = true;
overrideAllInputs(inputs);
overrideAllOutputs(outputs);
}
void InputModelTF::InputModelTFImpl::setPartialShape(ngraph::frontend::Place::Ptr place,
const ov::PartialShape& p_shape) {
castToTensorPlace(place)->set_partial_shape(p_shape);
}
ov::PartialShape InputModelTF::InputModelTFImpl::getPartialShape(ngraph::frontend::Place::Ptr place) const {
return castToTensorPlace(place)->get_partial_shape();
}
void InputModelTF::InputModelTFImpl::setElementType(ngraph::frontend::Place::Ptr place, const ov::element::Type& type) {
castToTensorPlace(place)->set_element_type(type);
}
void InputModelTF::InputModelTFImpl::setTensorValue(ngraph::frontend::Place::Ptr place, const void* value) {
m_graph_changed = true;
auto tensor_place = castToTensorPlace(place);
auto p_shape = tensor_place->get_partial_shape();
auto type = tensor_place->get_element_type();
auto constant = opset7::Constant::create(type, p_shape.to_shape(), value);
auto name = tensor_place->get_names()[0];
constant->set_friendly_name(name);
m_tensor_values[name] = constant;
}
InputModelTF::InputModelTF(const GraphIterator::Ptr& graph_iterator)
: _impl{std::make_shared<InputModelTFImpl>(graph_iterator, *this)} {}
std::vector<std::shared_ptr<OpPlaceTF>> InputModelTF::get_op_places() const {
return _impl->get_op_places();
}
std::map<std::string, std::shared_ptr<TensorPlaceTF>> InputModelTF::get_tensor_places() const {
return _impl->get_tensor_places();
}
std::map<std::string, Output<Node>> InputModelTF::get_tensor_values() const {
return _impl->get_tensor_values();
}
std::vector<ngraph::frontend::Place::Ptr> InputModelTF::get_inputs() const {
return _impl->getInputs();
}
std::vector<ngraph::frontend::Place::Ptr> InputModelTF::get_outputs() const {
return _impl->getOutputs();
}
ngraph::frontend::Place::Ptr InputModelTF::get_place_by_tensor_name(const std::string& tensorName) const {
return _impl->getPlaceByTensorName(tensorName);
}
void InputModelTF::override_all_outputs(const std::vector<ngraph::frontend::Place::Ptr>& outputs) {
_impl->overrideAllOutputs(outputs);
}
void InputModelTF::override_all_inputs(const std::vector<ngraph::frontend::Place::Ptr>& inputs) {
_impl->overrideAllInputs(inputs);
}
void InputModelTF::extract_subgraph(const std::vector<ngraph::frontend::Place::Ptr>& inputs,
const std::vector<ngraph::frontend::Place::Ptr>& outputs) {
_impl->extractSubgraph(inputs, outputs);
}
void InputModelTF::set_partial_shape(ngraph::frontend::Place::Ptr place, const ov::PartialShape& p_shape) {
_impl->setPartialShape(place, p_shape);
}
ov::PartialShape InputModelTF::get_partial_shape(ngraph::frontend::Place::Ptr place) const {
return _impl->getPartialShape(place);
}
void InputModelTF::set_element_type(ngraph::frontend::Place::Ptr place, const ov::element::Type& type) {
_impl->setElementType(place, type);
}
void InputModelTF::set_tensor_value(ngraph::frontend::Place::Ptr place, const void* value) {
_impl->setTensorValue(place, value);
}
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,39 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include "ngraph_conversions.hpp"
#include "utils.hpp"
namespace ov {
namespace frontend {
namespace tf {
void NHWCtoNCHW(const std::string& op_name, bool need_convert, ov::Output<ov::Node>& node) {
if (need_convert) {
auto rank = node.get_shape().size();
if (rank == 4) {
Transpose<0, 3, 1, 2>(node);
} else if (rank == 5) {
Transpose3D<0, 4, 1, 2, 3>(node);
}
SetTracingInfo(op_name, node);
}
}
void NCHWtoNHWC(const std::string& op_name, bool need_convert, ov::Output<ov::Node>& node) {
if (need_convert) {
auto rank = node.get_shape().size();
if (rank == 4) {
Transpose<0, 2, 3, 1>(node);
} else if (rank == 5) {
Transpose3D<0, 2, 3, 4, 1>(node);
}
SetTracingInfo(op_name, node);
}
}
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,96 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <openvino/opsets/opset8.hpp>
#include <string>
#include <tensorflow_frontend/utility.hpp>
#include "graph.pb.h"
#include "types.pb.h"
namespace ov {
namespace frontend {
namespace tf {
using ::tensorflow::DataType;
void TFTensorShapeToNGraphShape(const ::tensorflow::TensorShapeProto& tf_shape, ov::PartialShape* ng_shape);
template <size_t a, size_t b, size_t c, size_t d>
void Transpose(ov::Output<ov::Node>& node) {
static_assert(a < 4 && b < 4 && c < 4 && d < 4, "Number of dimensions cannot exceed 4");
static_assert(a != b && a != c && a != d && b != c && b != d && c != d, "Dimensions indices cannot be equal");
auto& s = node.get_shape();
ov::Shape reshaped_shape{s[a], s[b], s[c], s[d]};
ov::Shape transpose_order{a, b, c, d};
auto input_order =
std::make_shared<ov::opset8::Constant>(ov::element::u64, ov::Shape{transpose_order.size()}, transpose_order);
node = std::make_shared<ov::opset8::Transpose>(node, input_order);
}
template <size_t a, size_t b, size_t c, size_t d>
void Transpose(std::shared_ptr<ov::Node>& node) {
Transpose<a, b, c, d>(node->get_default_output());
}
template <size_t a, size_t b, size_t c, size_t d, size_t e>
void Transpose3D(ov::Output<ov::Node>& node) {
static_assert(a < 5 && b < 5 && c < 5 && d < 5 && e < 5, "Number of dimensions cannot exceed 5");
static_assert(a != b && a != c && a != d && a != e && b != c && b != d && b != e && c != d && c != e && d != e,
"Dimensions indices cannot be equal");
auto& s = node.get_shape();
ov::Shape reshaped_shape{s[a], s[b], s[c], s[d], s[e]};
ov::Shape transpose_order{a, b, c, d, e};
auto input_order =
std::make_shared<ov::opset8::Constant>(ov::element::u64, ov::Shape{transpose_order.size()}, transpose_order);
node = std::make_shared<ov::opset8::Transpose>(node, input_order);
}
template <size_t a, size_t b, size_t c, size_t d, size_t e>
void Transpose3D(std::shared_ptr<ov::Node>& node) {
Transpose3D<a, b, c, d, e>(node->get_default_output());
}
namespace detail {
template <typename T>
void NHWCtoHW(const std::vector<T>& src, std::vector<size_t>& dst) {
if (dst.size() >= 2) {
dst[0] = src[1];
dst[1] = src[2];
}
if (dst.size() >= 3) {
dst[2] = src[3];
}
}
template <typename T>
void NCHWtoHW(const std::vector<T>& src, std::vector<size_t>& dst) {
if (dst.size() >= 2) {
dst[0] = src[2];
dst[1] = src[3];
}
if (dst.size() >= 3) {
dst[2] = src[4];
}
}
} // namespace detail
void NHWCtoNCHW(const std::string& op_name, bool need_convert, ov::Output<ov::Node>& ng_input);
void NCHWtoNHWC(const std::string& op_name, bool need_convert, ov::Output<ov::Node>& ng_node);
template <typename T>
void NHWCtoHW(bool is_nhwc, const std::vector<T>& src, std::vector<size_t>& dst) {
if (is_nhwc) {
detail::NHWCtoHW(src, dst);
} else {
detail::NCHWtoHW(src, dst);
}
}
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,135 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <openvino/core/variant.hpp>
#include <tensorflow_frontend/exceptions.hpp>
#include <tensorflow_frontend/place.hpp>
#include <tensorflow_frontend/utility.hpp>
#include "tensor.pb.h"
#include "types.pb.h"
#define NGRAPH_VARIANT_DECLARATION(TYPE, info) \
template <> \
class VariantWrapper<TYPE> : public VariantImpl<TYPE> { \
public: \
OPENVINO_RTTI(info); \
VariantWrapper<TYPE>(const value_type& value) : VariantImpl<value_type>(value) {} \
}
namespace ov {
NGRAPH_VARIANT_DECLARATION(int32_t, "Variant::int32");
NGRAPH_VARIANT_DECLARATION(uint64_t, "Variant::uint64_t");
NGRAPH_VARIANT_DECLARATION(std::vector<int32_t>, "Variant::int32_vector");
NGRAPH_VARIANT_DECLARATION(float, "Variant::float");
NGRAPH_VARIANT_DECLARATION(std::vector<float>, "Variant::float_vector");
NGRAPH_VARIANT_DECLARATION(bool, "Variant::bool");
NGRAPH_VARIANT_DECLARATION(ov::element::Type, "Variant::ov_element_type");
NGRAPH_VARIANT_DECLARATION(std::vector<int64_t>, "Variant::int64_vector");
NGRAPH_VARIANT_DECLARATION(ov::PartialShape, "Variant::ngraph_PartialShape");
NGRAPH_VARIANT_DECLARATION(std::vector<std::string>, "Variant::string_vector");
NGRAPH_VARIANT_DECLARATION(::tensorflow::DataType, "Variant::DataType");
NGRAPH_VARIANT_DECLARATION(::tensorflow::TensorProto, "Variant::TensorProto");
} // namespace ov
namespace ov {
namespace frontend {
namespace tf {
using InPortName = size_t;
using OutPortName = size_t;
using NamedOutputs = std::map<OutPortName, OutputVector>;
using NamedInputs = std::map<InPortName, OutputVector>;
/// Keep necessary data for a single node in the original FW graph to facilitate
/// conversion process in the rules code.
class NodeContext {
const DecoderBase& m_decoder;
const NamedInputs& m_name_map;
public:
NodeContext(const DecoderBase& decoder, const NamedInputs& name_map) : m_decoder(decoder), m_name_map(name_map) {}
/// Returns node attribute by name. Returns 'def' value if attribute does not exist
template <typename T>
T get_attribute(const std::string& name, const T& def) const {
auto res = m_decoder.get_attribute(name, VariantWrapper<T>::get_type_info_static());
if (res) {
auto ret = std::dynamic_pointer_cast<VariantWrapper<T>>(res);
FRONT_END_GENERAL_CHECK(ret, "Attribute with name '", name, "' has invalid type");
return ret->get();
}
return def;
}
/// Returns node attribute by name
template <typename T>
T get_attribute(const std::string& name) const {
auto res = m_decoder.get_attribute(name, VariantWrapper<T>::get_type_info_static());
FRONT_END_GENERAL_CHECK(res, "Attribute with name '", name, "' does not exist");
auto ret = std::dynamic_pointer_cast<VariantWrapper<T>>(res);
FRONT_END_GENERAL_CHECK(ret, "Attribute with name '", name, "' has invalid type");
return ret->get();
}
/// Check if an attribute of a given name exists
template <typename T>
bool has_attribute(const std::string& name) const {
return m_decoder.get_attribute(name, VariantWrapper<T>::get_type_info_static()) != nullptr;
}
/// Detects if there is at least one input attached with a given name
bool has_ng_input(const size_t& port_index) const {
auto found = m_name_map.find(port_index);
if (found != m_name_map.end())
return !found->second.empty();
return false;
}
/// Returns exactly one input with a given name; throws if there is no inputs or
/// there are more than one input
Output<Node> get_ng_input(const size_t& port_index) const {
FRONT_END_GENERAL_CHECK(m_name_map.at(port_index).size() == 1);
return m_name_map.at(port_index).at(0);
}
/// Returns all inputs with a given name
OutputVector get_ng_inputs(const size_t& port_index) const {
return m_name_map.at(port_index);
}
/// Returns all inputs in order they appear in map. This is used for FrameworkNode
/// creation
OutputVector get_all_ng_inputs() const {
OutputVector res;
for (const auto& entry : m_name_map) {
res.insert(res.end(), entry.second.begin(), entry.second.end());
}
return res;
}
/// Get a number of inputs
size_t get_ng_input_size() const {
return m_name_map.size();
}
/// Get operation type
std::string get_op_type() const {
return m_decoder.get_op_type();
}
/// Get a node name
std::string get_name() const {
return m_decoder.get_op_name();
}
/// Get a decoder
const DecoderBase* get_decoder() const {
return &m_decoder;
}
};
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,33 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <numeric>
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateAddNOp(const NodeContext& node) {
OutputVector ng_arg_vec = node.get_all_ng_inputs();
auto ng_addn = std::accumulate(std::next(ng_arg_vec.begin()),
ng_arg_vec.end(),
ng_arg_vec.at(0),
[&node](Output<Node> a, Output<Node> b) {
return ConstructNgNode<Add>(node.get_name(), a, b);
}); // accumulation: start with
// first element. default op is
// addition
return {ng_addn};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,62 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateArgMinMax(const NodeContext& node, std::string mode) {
Output<Node> ng_input = node.get_ng_input(0);
std::vector<int64_t> tf_dim;
GetStaticInputVector(node, 1, &tf_dim);
Shape input_shape = ng_input.get_shape();
size_t input_rank = input_shape.size();
if (tf_dim.size() != 1) {
throw errors::InvalidArgument("ArgMax Op: dimension must be scalar, operates on a single axis");
}
// If input dimension is negative, make it positive
if (tf_dim[0] < 0) {
NGRAPH_VLOG(3) << "Input dimension is negative, make it positive " << tf_dim[0];
tf_dim[0] = (int64_t)input_rank + tf_dim[0];
}
NGRAPH_VLOG(3) << "Axis along which to compute " << tf_dim[0];
size_t k_axis = tf_dim[0];
auto ng_et = node.get_attribute<element::Type>("output_type");
auto ng_k = ConstructNgNode<Constant>(node.get_name(), element::i64, Shape{}, std::vector<int64_t>({1}));
std::string sort = "none";
auto ng_topk = std::make_shared<TopK>(ng_input, ng_k, k_axis, mode, sort, ng_et);
auto ng_indices = ng_topk->output(1);
int axis = ng_topk->get_axis();
auto axis_to_remove =
ConstructNgNode<Constant>(node.get_name(), element::i64, Shape{1}, std::vector<int64_t>({axis}));
auto reshaped_indices = ConstructNgNode<Squeeze>(node.get_name(), ng_indices, axis_to_remove);
SetTracingInfo(node.get_name(), reshaped_indices);
return {reshaped_indices};
}
OutputVector TranslateArgMaxOp(const NodeContext& node) {
return (TranslateArgMinMax(node, "max"));
}
OutputVector TranslateArgMinOp(const NodeContext& node) {
return (TranslateArgMinMax(node, "min"));
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,69 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateAvgPoolOp(const NodeContext& node) {
Output<Node> ng_input = node.get_ng_input(0);
auto tf_strides = node.get_attribute<std::vector<int32_t>>("strides");
auto tf_ksize = node.get_attribute<std::vector<int32_t>>("ksize");
auto tf_padding_type = node.get_attribute<std::string>("padding");
auto tf_data_format = node.get_attribute<std::string>("data_format");
if (tf_data_format != "NHWC" && tf_data_format != "NCHW") {
throw errors::InvalidArgument("AvgPool data format is neither NHWC nor NCHW");
}
bool is_nhwc = (tf_data_format == "NHWC");
Strides ng_strides(2);
Shape ng_image_shape(2);
Shape ng_kernel_shape(2);
NHWCtoHW(is_nhwc, tf_strides, ng_strides);
NHWCtoHW(is_nhwc, ng_input.get_shape(), ng_image_shape);
NHWCtoHW(is_nhwc, tf_ksize, ng_kernel_shape);
NHWCtoNCHW(node.get_name(), is_nhwc, ng_input);
CoordinateDiff padding_below;
CoordinateDiff padding_above;
Shape ng_dilations{1, 1};
MakePadding(tf_padding_type,
ng_image_shape,
ng_kernel_shape,
ng_strides,
ng_dilations,
padding_below,
padding_above);
// TODO: remove this once nGraph supports negative padding
// (CoordinateDiff) for AvgPool
Shape ng_padding_below(padding_below.begin(), padding_below.end());
Shape ng_padding_above(padding_above.begin(), padding_above.end());
Output<Node> ng_avgpool = ConstructNgNode<AvgPool>(node.get_name(),
ng_input,
ng_strides,
ng_padding_below,
ng_padding_above,
ng_kernel_shape,
true,
ov::op::RoundingType::FLOOR);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_avgpool);
return {ng_avgpool};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,55 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateBiasAddOp(const NodeContext& node) {
Output<Node> ng_input = node.get_ng_input(0), ng_bias = node.get_ng_input(1);
std::string tf_data_format = node.get_attribute<std::string>("data_format", "NHWC");
if (tf_data_format != "NHWC" && tf_data_format != "NCHW") {
throw errors::InvalidArgument("BiasAdd data format is neither NHWC nor NCHW");
}
auto ng_input_shape = ng_input.get_shape();
auto ng_bias_shape = ng_bias.get_shape();
if (ng_bias_shape.size() != 1) {
throw errors::InvalidArgument("Bias argument to BiasAdd does not have one dimension");
}
// We'll choose reshape over broadcast
// Reshape the bias to (1, C, 1, ...) if input is channels-first.
Output<Node> ng_bias_reshaped = ng_bias;
if (tf_data_format == "NCHW") {
auto channel_dim = ng_input_shape[1];
std::vector<int64_t> target_shape(ng_input_shape.size());
for (int64_t i = 0; i < ng_input_shape.size(); i++) {
if (i == 1) {
target_shape[i] = channel_dim;
} else {
target_shape[i] = 1;
}
}
auto target_shape_node = make_shared<Constant>(element::i64, Shape{ng_input_shape.size()}, target_shape);
ng_bias_reshaped = ConstructNgNode<Reshape>(node.get_name(), ng_bias, target_shape_node, false);
}
Output<Node> ng_add = ConstructNgNode<Add>(node.get_name(), ng_input, ng_bias_reshaped);
return {ng_add};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,99 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
// Helper function to translate a binary op
// Parameters:
//
// TFNodeDecoder* op - TF op being translated. Must have only two
// inputs.
// const std::vector<const ov::frontend::tf::detail::TensorWrapper*>& static_input_map - the static input
// map Builder::OpMap& ng_op_map - The TF-to-nGraph op map. std::function<Output<Node>(Output<Node>,
// Output<Node>)>
// create_binary_op - Function to construct the graph implementing
// the binary op, given the 2 ng_inputs to the
// binaryop
// Example Usage:
//
// if (op->type_string() == "SquaredDifference") {
// TF_RETURN_IF_ERROR(TranslateBinaryOp(op, ng_op_map,
// [](Output<Node> ng_input1, Output<Node>
// ng_input2) {
// auto ng_diff = Output<Subtract>(input1,
// input2);
// return Output<Multiply>(ng_diff,ng_diff);
// }));
// }
//
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateBinaryOp(const NodeContext& node,
std::function<Output<Node>(Output<Node>&, Output<Node>&)> create_binary_op) {
Output<Node> ng_lhs = node.get_ng_input(0), ng_rhs = node.get_ng_input(1);
auto ng_node = create_binary_op(ng_lhs, ng_rhs);
// TODO do we need it?
/* if (ng_node != ng_lhs && ng_node != ng_rhs) {
Builder::SetTracingInfo(node.get_name(), ng_node);
}*/
return {ng_node};
}
OutputVector TranslateFloorDivOp(const NodeContext& node) {
auto floordiv_fn = [&node](Output<Node> x, Output<Node> y) {
return ConstructNgNode<Floor>(node.get_name(), ConstructNgNode<Divide>(node.get_name(), x, y));
};
return TranslateBinaryOp(node, floordiv_fn);
}
// Helper function to translate a binary op in cases where there is a one-to-one
// mapping from TensorFlow ops to nGraph ops.
//
// Example usage:
//
// if (n->type_string == "Add") {
// TF_RETURN_IF_ERROR(TranslateBinaryOp<Add>(op,
// static_input_map,
// ng_op_map));
// }
//
template <typename T>
OutputVector TranslateBinaryOp(const NodeContext& node) {
return TranslateBinaryOp(node, [&node](Output<Node>& ng_lhs, Output<Node>& ng_rhs) {
return ConstructNgNode<T>(node.get_name(), ng_lhs, ng_rhs);
});
}
template OutputVector TranslateBinaryOp<Add>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Equal>(const NodeContext& node);
template OutputVector TranslateBinaryOp<FloorMod>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Greater>(const NodeContext& node);
template OutputVector TranslateBinaryOp<GreaterEqual>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Less>(const NodeContext& node);
template OutputVector TranslateBinaryOp<LessEqual>(const NodeContext& node);
template OutputVector TranslateBinaryOp<LogicalAnd>(const NodeContext& node);
template OutputVector TranslateBinaryOp<LogicalOr>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Maximum>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Minimum>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Multiply>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Mod>(const NodeContext& node);
template OutputVector TranslateBinaryOp<NotEqual>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Power>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Divide>(const NodeContext& node);
template OutputVector TranslateBinaryOp<SquaredDifference>(const NodeContext& node);
template OutputVector TranslateBinaryOp<Subtract>(const NodeContext& node);
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,26 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateCastOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0);
auto ng_et = node.get_attribute<element::Type>("DstT");
return {ConstructNgNode<Convert>(node.get_name(), ng_input, ng_et)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,44 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov;
using namespace opset8;
using namespace ov::frontend;
using namespace frontend::tf::detail;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateConcatV2Op(const NodeContext& node) {
ValidateInputCountMin(node, 2);
std::vector<int64_t> tf_concat_axis_vec;
GetStaticInputVector(node, node.get_ng_input_size() - 1, &tf_concat_axis_vec);
int64_t concat_axis = tf_concat_axis_vec[0];
if (concat_axis < 0) {
auto ng_first_arg = node.get_ng_input(0);
concat_axis += int64_t(ng_first_arg.get_shape().size());
}
OutputVector ng_args;
for (int i = 0; i < node.get_ng_input_size() - 1; i++) {
Output<Node> ng_arg = node.get_ng_input(i);
ng_args.push_back(ng_arg);
}
return {ConstructNgNode<Concat>(node.get_name(), ng_args, size_t(concat_axis))};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,68 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
namespace {
using ConstMap = std::map<ov::element::Type,
std::pair<std::function<Status(const NodeContext&, ov::element::Type, ov::Output<ov::Node>&)>,
const ov::element::Type>>;
const ConstMap& TF_NGRAPH_CONST_MAP() {
static const ConstMap the_map = {
{ov::element::f32, make_pair(MakeConstOp<float>, ov::element::f32)},
{ov::element::f64, make_pair(MakeConstOp<double>, ov::element::f64)},
{ov::element::i8, make_pair(MakeConstOp<int8_t>, ov::element::i8)},
{ov::element::i16, make_pair(MakeConstOp<int16_t>, ov::element::i16)},
#if 0
{DataType::DT_QINT8, make_pair(MakeConstOp<qint8>, ov::element::i8)},
{DataType::DT_QUINT8, make_pair(MakeConstOp<quint8>, ov::element::u8)},
{DataType::DT_QUINT16, make_pair(MakeConstOp<quint16>, ov::element::u16)},
#endif
{ov::element::i32, make_pair(MakeConstOp<int32_t>, ov::element::i32)},
{ov::element::i64, make_pair(MakeConstOp<int64_t>, ov::element::i64)},
{ov::element::u8, make_pair(MakeConstOp<uint8_t>, ov::element::u8)},
{ov::element::u16, make_pair(MakeConstOp<uint16_t>, ov::element::u16)},
{ov::element::boolean, make_pair(MakeConstOp<bool, char>, ov::element::boolean)}
};
return the_map;
}
} // namespace
OutputVector TranslateConstOp(const NodeContext& node) {
auto dt = node.get_attribute<ov::element::Type>("dtype");
Output<Node> ng_node;
// For some reason the following do not work (no specialization of
// tensorflow::checkpoint::SavedTypeTraits...)
// case DataType::DT_UINT32:
// TF_RETURN_IF_ERROR(MakeConstOp<uint32>(op, element::u32,
// &ng_node));
// break;
// case DataType::DT_UINT64:
// TF_RETURN_IF_ERROR(MakeConstOp<uint64>(op, element::u64,
// &ng_node));
// break;
try {
const auto& func_param = TF_NGRAPH_CONST_MAP().at(dt);
TF_RETURN_IF_ERROR(func_param.first(node, func_param.second, ng_node));
} catch (const std::out_of_range&) {
throw errors::Unimplemented("Failed to translate Constant with target ngraph type:" + dt.get_type_name());
}
return {ng_node};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,76 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateConv2DOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0), ng_filter = node.get_ng_input(1);
auto tf_strides = node.get_attribute<std::vector<int32_t>>("strides");
auto tf_dilations = node.get_attribute<std::vector<int32_t>>("dilations");
auto tf_padding_type = node.get_attribute<std::string>("padding");
auto tf_data_format = node.get_attribute<std::string>("data_format");
if (tf_data_format != "NHWC" && tf_data_format != "NCHW") {
throw errors::InvalidArgument("Conv2D data format is neither NHWC nor NCHW");
}
bool is_nhwc = (tf_data_format == "NHWC");
// TF Kernel Test Checks
// Strides in the batch and depth dimension is not supported
if (tf_strides[0] != 1 || tf_strides[is_nhwc ? 3 : 1] != 1) {
throw errors::InvalidArgument("Strides in batch and depth dimensions is not supported: " + node.get_op_type());
}
Strides ng_strides(2);
Strides ng_dilations(2);
Shape ng_image_shape(2);
Shape ng_kernel_shape(2);
NHWCtoHW(is_nhwc, tf_strides, ng_strides);
NHWCtoHW(is_nhwc, ng_input.get_shape(), ng_image_shape);
NHWCtoHW(is_nhwc, tf_dilations, ng_dilations);
NHWCtoNCHW(node.get_name(), is_nhwc, ng_input);
auto& ng_filter_shape = ng_filter.get_shape();
ng_kernel_shape[0] = ng_filter_shape[0];
ng_kernel_shape[1] = ng_filter_shape[1];
Transpose<3, 2, 0, 1>(ng_filter);
SetTracingInfo(node.get_name(), ng_filter);
CoordinateDiff ng_padding_below;
CoordinateDiff ng_padding_above;
MakePadding(tf_padding_type,
ng_image_shape,
ng_kernel_shape,
ng_strides,
ng_dilations,
ng_padding_below,
ng_padding_above);
Output<Node> ng_conv = ConstructNgNode<Convolution>(node.get_name(),
ng_input,
ng_filter,
ng_strides,
ng_padding_below,
ng_padding_above,
ng_dilations);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_conv);
return {ng_conv};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,98 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateConv2DBackpropInputOp(const NodeContext& node) {
auto ng_filter = node.get_ng_input(1), ng_out_backprop = node.get_ng_input(2);
// TODO: refactor me to be less redundant with other convolution ops
auto tf_strides = node.get_attribute<std::vector<int32_t>>("strides");
auto tf_dilations = node.get_attribute<std::vector<int32_t>>("dilations");
auto tf_padding_type = node.get_attribute<std::string>("padding");
auto tf_data_format = node.get_attribute<std::string>("data_format");
if (tf_data_format != "NHWC" && tf_data_format != "NCHW") {
throw errors::InvalidArgument("Conv2DBackpropInput data format is neither NHWC nor NCHW: %s" + tf_data_format);
}
std::vector<int64_t> tf_input_sizes;
GetStaticInputVector(node, 0, &tf_input_sizes);
if (std::any_of(tf_input_sizes.begin(), tf_input_sizes.end(), [](int32_t size) {
return size <= 0;
})) {
throw errors::InvalidArgument("Conv2DBackpropInput input sizes must be positive integers");
}
bool is_nhwc = (tf_data_format == "NHWC");
Strides ng_strides(2);
Strides ng_dilations(2);
Shape ng_image_shape(2);
Shape ng_kernel_shape(2);
Shape ng_batch_shape(4);
NHWCtoHW(is_nhwc, tf_strides, ng_strides);
NHWCtoHW(is_nhwc, tf_dilations, ng_dilations);
NHWCtoHW(is_nhwc, tf_input_sizes, ng_image_shape);
NHWCtoNCHW(node.get_name(), is_nhwc, ng_out_backprop);
if (is_nhwc) {
ng_batch_shape = {static_cast<unsigned long>(tf_input_sizes[0]),
static_cast<unsigned long>(tf_input_sizes[3]),
static_cast<unsigned long>(tf_input_sizes[1]),
static_cast<unsigned long>(tf_input_sizes[2])};
} else {
ng_batch_shape = {static_cast<unsigned long>(tf_input_sizes[0]),
static_cast<unsigned long>(tf_input_sizes[1]),
static_cast<unsigned long>(tf_input_sizes[2]),
static_cast<unsigned long>(tf_input_sizes[3])};
}
auto& ng_filter_shape = ng_filter.get_shape();
ng_kernel_shape[0] = ng_filter_shape[0];
ng_kernel_shape[1] = ng_filter_shape[1];
Transpose<3, 2, 0, 1>(ng_filter);
SetTracingInfo(node.get_name(), ng_filter);
CoordinateDiff ng_padding_below;
CoordinateDiff ng_padding_above;
MakePadding(tf_padding_type,
ng_image_shape,
ng_kernel_shape,
ng_strides,
ng_dilations,
ng_padding_below,
ng_padding_above);
auto ng_output_shape = ConstructNgNode<Constant>(node.get_name(),
element::i64,
Shape{ng_batch_shape.size() - 2},
vector<size_t>(ng_batch_shape.begin() + 2, ng_batch_shape.end()));
auto ng_data = ConstructNgNode<ConvolutionBackpropData>(node.get_name(),
ng_out_backprop,
ng_filter,
ng_output_shape,
ng_strides,
ng_padding_below,
ng_padding_above,
ng_dilations);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_data);
return {ng_data};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,81 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
// Translate Conv3D Op
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateConv3DOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0), ng_filter = node.get_ng_input(1);
auto tf_strides = node.get_attribute<std::vector<int32_t>>("strides");
auto tf_dilations = node.get_attribute<std::vector<int32_t>>("dilations");
auto tf_padding_type = node.get_attribute<std::string>("padding");
auto tf_data_format = node.get_attribute<std::string>("data_format");
if (tf_data_format != "NDHWC" && tf_data_format != "NCDHW") {
throw errors::InvalidArgument("Conv3D data format is neither NDHWC nor NCDHW");
}
bool is_ndhwc = (tf_data_format == "NDHWC");
// TODO: in 3D
// TF Kernel Test Checks
// // Strides in the batch and depth dimension is not supported
// if (tf_strides[0] != 1 || tf_strides[is_nhwc ? 3 : 1] != 1) {
// return errors::InvalidArgument(
// "Strides in batch and depth dimensions is not supported: ",
// op->type_string());
// }
Strides ng_strides(3);
Strides ng_dilations(3);
Shape ng_image_shape(3);
Shape ng_kernel_shape(3);
NHWCtoHW(is_ndhwc, tf_strides, ng_strides);
NHWCtoHW(is_ndhwc, ng_input.get_shape(), ng_image_shape);
NHWCtoHW(is_ndhwc, tf_dilations, ng_dilations);
NHWCtoNCHW(node.get_name(), is_ndhwc, ng_input);
auto& ng_filter_shape = ng_filter.get_shape();
ng_kernel_shape[0] = ng_filter_shape[0];
ng_kernel_shape[1] = ng_filter_shape[1];
ng_kernel_shape[2] = ng_filter_shape[2];
Transpose3D<4, 3, 0, 1, 2>(ng_filter);
SetTracingInfo(node.get_name(), ng_filter);
CoordinateDiff ng_padding_below;
CoordinateDiff ng_padding_above;
MakePadding(tf_padding_type,
ng_image_shape,
ng_kernel_shape,
ng_strides,
ng_dilations,
ng_padding_below,
ng_padding_above);
Output<Node> ng_conv = ConstructNgNode<Convolution>(node.get_name(),
ng_input,
ng_filter,
ng_strides,
ng_padding_below,
ng_padding_above,
ng_dilations);
NCHWtoNHWC(node.get_name(), is_ndhwc, ng_conv);
return {ng_conv};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,25 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateCumsumOp(const NodeContext& node) {
auto ng_x = node.get_ng_input(0), ng_axis = node.get_ng_input(1);
auto exclusive = node.get_attribute<bool>("exclusive"), reverse = node.get_attribute<bool>("reverse");
return {ConstructNgNode<CumSum>(node.get_name(), ng_x, ng_axis, exclusive, reverse)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,39 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
// Translate DepthToSpace op
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateDepthToSpaceOp(const NodeContext& node) {
Output<Node> ng_input = node.get_ng_input(0);
// Get the attributes
auto block_size = node.get_attribute<int64_t>("block_size");
std::string tf_data_format = node.get_attribute<std::string>("data_format");
if (tf_data_format != "NHWC" && tf_data_format != "NCHW") {
throw errors::InvalidArgument("DepthToSpace data format is neither NHWC nor NCHW");
}
bool is_nhwc = (tf_data_format == "NHWC");
NHWCtoNCHW(node.get_name(), is_nhwc, ng_input);
auto ng_mode = DepthToSpace::DepthToSpaceMode::BLOCKS_FIRST;
Output<Node> depth_to_space = ConstructNgNode<DepthToSpace>(node.get_name(), ng_input, ng_mode, block_size);
NCHWtoNHWC(node.get_name(), is_nhwc, depth_to_space);
return {depth_to_space};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,80 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateDepthwiseConv2dNativeOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0), ng_filter = node.get_ng_input(1);
auto tf_strides = node.get_attribute<std::vector<int32_t>>("strides");
auto tf_dilations = node.get_attribute<std::vector<int32_t>>("dilations");
auto tf_padding_type = node.get_attribute<std::string>("padding");
auto tf_data_format = node.get_attribute<std::string>("data_format");
if (tf_data_format != "NHWC" && tf_data_format != "NCHW") {
throw errors::InvalidArgument("DepthwiseConv2D data format is neither NHWC nor NCHW");
}
bool is_nhwc = (tf_data_format == "NHWC");
Strides ng_strides(2);
Strides ng_dilations(2);
Shape ng_image_shape(2);
Shape ng_kernel_shape(2);
NHWCtoHW(is_nhwc, ng_input.get_shape(), ng_image_shape);
NHWCtoHW(is_nhwc, tf_strides, ng_strides);
NHWCtoHW(is_nhwc, tf_dilations, ng_dilations);
NHWCtoNCHW(node.get_name(), is_nhwc, ng_input);
auto& ng_filter_shape = ng_filter.get_shape();
ng_kernel_shape[0] = ng_filter_shape[0];
ng_kernel_shape[1] = ng_filter_shape[1];
CoordinateDiff ng_padding_below;
CoordinateDiff ng_padding_above;
MakePadding(tf_padding_type,
ng_image_shape,
ng_kernel_shape,
ng_strides,
ng_dilations,
ng_padding_below,
ng_padding_above);
// H W I M -> H W I 1 M
auto filter_shape = ConstructNgNode<Constant>(
node.get_name(),
element::u64,
Shape{5},
ov::Shape{ng_filter_shape[0], ng_filter_shape[1], ng_filter_shape[2], 1, ng_filter_shape[3]});
auto reshaped_filter = ConstructNgNode<Reshape>(node.get_name(), ng_filter, filter_shape, false);
// H W I 1 M -> I M 1 H W
auto order = ConstructNgNode<Constant>(node.get_name(), element::i64, Shape{5}, vector<int64_t>{2, 4, 3, 0, 1});
auto transposed_filter = ConstructNgNode<ov::opset8::Transpose>(node.get_name(), reshaped_filter, order);
auto ng_conv = ConstructNgNode<GroupConvolution>(node.get_name(),
ng_input,
transposed_filter,
ng_strides,
ng_padding_below,
ng_padding_above,
ng_dilations);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_conv);
return {ng_conv};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,27 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
#include "node_context.hpp"
using namespace std;
using namespace ov;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateEluOp(const NodeContext& node) {
auto input = node.get_ng_input(0);
auto alpha = 1.0; // node.get_attribute<float>("alpha");
return {ConstructNgNode<Elu>(node.get_name(), input, alpha)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,26 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateExpandDimsOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0);
std::vector<int64_t> dims;
GetStaticInputVector(node, 1, &dims);
auto ng_dims = ConstructNgNode<Constant>(node.get_name(), element::i64, ov::Shape{dims.size()}, dims);
return {ConstructNgNode<Unsqueeze>(node.get_name(), ng_input, ng_dims)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,23 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateFillOp(const NodeContext& node) {
auto ng_dims = node.get_ng_input(0), ng_value = node.get_ng_input(1);
return {ConstructNgNode<Broadcast>(node.get_name(), ng_value, ng_dims)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,57 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateFusedBatchNormOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0), ng_scale = node.get_ng_input(1), ng_offset = node.get_ng_input(2),
ng_mean = node.get_ng_input(3), ng_variance = node.get_ng_input(4);
bool is_v3 = node.get_op_type() == "FusedBatchNormV3";
auto tf_data_format = node.get_attribute<std::string>("data_format");
if (tf_data_format != "NHWC" && tf_data_format != "NCHW") {
throw errors::InvalidArgument("Conv2D data format is neither NHWC nor NCHW");
}
bool is_nhwc = (tf_data_format == "NHWC");
NGRAPH_VLOG(3) << "data_format: " << tf_data_format;
auto tf_epsilon = node.get_attribute<float>("epsilon", 0.0001); // TODO: where does 0.0001 come from?
NGRAPH_VLOG(3) << "epsilon: " << tf_epsilon;
NHWCtoNCHW(node.get_name(), is_nhwc, ng_input);
auto ng_batch_norm = ConstructNgNode<BatchNormInference>(node.get_name(),
ng_input,
ng_scale,
ng_offset,
ng_mean,
ng_variance,
tf_epsilon);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_batch_norm);
// TODO: Why are there so many? Is it correct?
OutputVector result = {ng_batch_norm, ng_mean, ng_variance, ng_mean, ng_variance};
if (is_v3) {
// FusedBatchNormV3 has 6 outputs
result.push_back(ng_mean); // reserve_space_3
}
return result;
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,148 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateFusedConv2DOp(const NodeContext& node) {
auto num_args = node.get_attribute<int>("num_args");
auto fused_ops = node.get_attribute<std::vector<string>>("fused_ops");
auto tf_data_format = node.get_attribute<std::string>("data_format");
bool is_nhwc = (tf_data_format == "NHWC");
auto CreateNgConv = [&](Output<Node>& ng_input, Output<Node>& ng_filter) {
auto tf_strides = node.get_attribute<std::vector<int32_t>>("strides");
auto tf_dilations = node.get_attribute<std::vector<int32_t>>("dilations");
auto tf_padding_type = node.get_attribute<std::string>("padding");
if (tf_data_format != "NHWC" && tf_data_format != "NCHW") {
throw errors::InvalidArgument("Conv2D data format is neither NHWC nor NCHW");
}
// TF Kernel Test Checks
// Strides in the batch and depth dimension is not supported
if (tf_strides[0] != 1 || tf_strides[is_nhwc ? 3 : 1] != 1) {
throw errors::InvalidArgument("Strides in batch and depth dimensions is not supported: " +
node.get_op_type());
}
Strides ng_strides(2);
Strides ng_dilations(2);
Shape ng_image_shape(2);
Shape ng_kernel_shape(2);
NHWCtoHW(is_nhwc, tf_strides, ng_strides);
NHWCtoHW(is_nhwc, ng_input.get_shape(), ng_image_shape);
NHWCtoHW(is_nhwc, tf_dilations, ng_dilations);
NHWCtoNCHW(node.get_name(), is_nhwc, ng_input);
auto& ng_filter_shape = ng_filter.get_shape();
ng_kernel_shape[0] = ng_filter_shape[0];
ng_kernel_shape[1] = ng_filter_shape[1];
Transpose<3, 2, 0, 1>(ng_filter);
SetTracingInfo(node.get_name(), ng_filter);
CoordinateDiff ng_padding_below;
CoordinateDiff ng_padding_above;
MakePadding(tf_padding_type,
ng_image_shape,
ng_kernel_shape,
ng_strides,
ng_dilations,
ng_padding_below,
ng_padding_above);
return ConstructNgNode<Convolution>(node.get_name() + "_FusedConv2D_Conv",
ng_input,
ng_filter,
ng_strides,
ng_padding_below,
ng_padding_above,
ng_dilations);
};
if (VecStrCmp(fused_ops, {"BiasAdd"}) || VecStrCmp(fused_ops, {"BiasAdd", "Relu"}) ||
VecStrCmp(fused_ops, {"BiasAdd", "Relu6"})) {
if (num_args != 1) {
throw errors::InvalidArgument("FusedConv2DBiasAdd has incompatible num_args");
}
auto ng_input = node.get_ng_input(0), ng_filter = node.get_ng_input(1), ng_bias = node.get_ng_input(2),
ng_conv = CreateNgConv(ng_input, ng_filter);
auto ng_conv_shape = ng_conv.get_shape();
auto ng_bias_shape = ng_bias.get_shape();
if (ng_bias_shape.size() != 1) {
throw errors::InvalidArgument("Bias argument to BiasAdd does not have one dimension");
}
std::vector<size_t> reshape_pattern_values(ng_conv_shape.size(), 1U);
reshape_pattern_values[1] = ng_bias.get_shape().front();
auto reshape_pattern =
make_shared<Constant>(element::u64, Shape{reshape_pattern_values.size()}, reshape_pattern_values);
auto ng_bias_reshaped = ConstructNgNode<Reshape>(node.get_name(), ng_bias, reshape_pattern, false);
auto ng_add = ConstructNgNode<Add>(node.get_name() + "_FusedConv2D_BiasAdd", ng_conv, ng_bias_reshaped);
if (VecStrCmp(fused_ops, {"BiasAdd", "Relu"})) {
auto ng_relu = ConstructNgNode<Relu>(node.get_name() + "_FusedConv2D_Relu", ng_add);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_relu);
return {ng_relu};
} else if (VecStrCmp(fused_ops, {"BiasAdd", "Relu6"})) {
auto ng_relu6 = ConstructNgNode<Clamp>(node.get_name() + "_FusedConv2D_Relu6", ng_add, 0, 6);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_relu6);
return {ng_relu6};
} else {
NCHWtoNHWC(node.get_name(), is_nhwc, ng_add);
return {ng_add};
}
} else if (VecStrCmp(fused_ops, {"FusedBatchNorm"}) || VecStrCmp(fused_ops, {"FusedBatchNorm", "Relu"}) ||
VecStrCmp(fused_ops, {"FusedBatchNorm", "Relu6"})) {
if (num_args != 4) {
throw errors::InvalidArgument("FusedConv2D with FusedBatchNorm has incompatible num_args");
}
auto ng_input = node.get_ng_input(0), ng_filter = node.get_ng_input(1), ng_scale = node.get_ng_input(2),
ng_offset = node.get_ng_input(3), ng_mean = node.get_ng_input(4), ng_variance = node.get_ng_input(5),
ng_conv = CreateNgConv(ng_input, ng_filter);
auto tf_epsilon = node.get_attribute<float>("epsilon");
auto ng_batch_norm = ConstructNgNode<BatchNormInference>(node.get_name() + "_FusedConv2D_BatchNorm",
ng_conv,
ng_scale,
ng_offset,
ng_mean,
ng_variance,
tf_epsilon);
if (VecStrCmp(fused_ops, {"FusedBatchNorm", "Relu"})) {
auto ng_relu = ConstructNgNode<Relu>(node.get_name() + "_FusedConv2D_BatchNormRelu", ng_batch_norm);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_relu);
return {ng_relu};
} else if (VecStrCmp(fused_ops, {"FusedBatchNorm", "Relu6"})) {
auto ng_relu6 = ConstructNgNode<Clamp>(node.get_name() + "_FusedConv2D_BatchNormRelu", ng_batch_norm, 0, 6);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_relu6);
return {ng_relu6};
} else {
NCHWtoNHWC(node.get_name(), is_nhwc, ng_batch_norm);
return {ng_batch_norm};
}
} else {
FRONT_END_THROW("Unsupported _FusedConv2D ");
}
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,54 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateFusedMatMulOp(const NodeContext& node) {
// auto num_args = node.get_attribute<int>("num_args"); // TODO: it is unused but why?
auto fused_ops = node.get_attribute<std::vector<string>>("fused_ops");
// Transpose arguments if requested.
auto transpose_a = node.get_attribute<bool>("transpose_a", false);
auto transpose_b = node.get_attribute<bool>("transpose_b", false);
auto ng_lhs = node.get_ng_input(0), ng_rhs = node.get_ng_input(1), ng_bias = node.get_ng_input(2);
Output<Node> ng_matmul = ConstructNgNode<MatMul>(node.get_name(), ng_lhs, ng_rhs, transpose_a, transpose_b);
auto ng_matmul_shape = ng_matmul.get_shape();
auto ng_bias_shape = ng_bias.get_shape();
if (ng_bias_shape.size() != 1) {
throw errors::InvalidArgument("Bias argument to BiasAdd does not have one dimension");
}
auto ng_add = ConstructNgNode<Add>(node.get_name(), ng_matmul, ng_bias);
if (fused_ops.size() == 1) { // Only fusing BiasAdd
return {ng_add};
} else if (fused_ops.size() == 2) { // Also has activation
if (fused_ops[1] == "Relu") {
return {ConstructNgNode<Relu>(node.get_name(), ng_add)};
} else if (fused_ops[1] == "Relu6") {
return {ConstructNgNode<Clamp>(node.get_name(), ng_add, 0, 6)};
} else {
throw errors::Internal("Expected activation to be Relu or Relu6 but got " + fused_ops[1]);
}
} else {
// Adding this here to catch future changes in _FusedMatMul
throw errors::Internal("Unsupported combination");
}
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,66 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
// See .../tensorflow/include/tensorflow/cc/ops/array_ops.h
// and .../openvino/ngraph/core/include/ngraph/op/gather.hpp
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateGatherOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0), ng_input_indices = node.get_ng_input(1);
auto ng_axis = ConstructNgNode<Constant>(node.get_name(), element::i64, Shape{}, 0);
auto gather_op = ConstructNgNode<Gather>(node.get_name(), ng_input, ng_input_indices, ng_axis);
return {gather_op};
}
OutputVector TranslateGatherV2Op(const NodeContext& node) {
auto ng_input = node.get_ng_input(0), ng_input_coords = node.get_ng_input(1);
std::vector<int64_t> tf_axis;
GetStaticInputVector(node, 2, &tf_axis);
if (tf_axis.size() > 1) {
std::ostringstream buf;
buf << "Found axis in GatherV2 op (" << node.get_name() << ") translation to be non scalar, of size "
<< tf_axis.size();
throw errors::Internal(buf.str());
}
// Negative axis is supported. Accounting for that
auto ng_input_shape = ng_input.get_shape();
size_t ng_input_rank = ng_input_shape.size();
int axis;
if (tf_axis[0] >= 0) {
axis = tf_axis[0];
} else {
axis = tf_axis[0] + ng_input_rank;
}
if (axis < 0 || axis >= ng_input_rank) {
std:
ostringstream buf;
buf << "Expected axis in the range [-" << ng_input_rank << ", " << ng_input_rank << "), but got " << tf_axis[0];
throw errors::InvalidArgument(buf.str());
}
auto ng_axis = ConstructNgNode<Constant>(node.get_name(), element::i64, Shape{tf_axis.size()}, tf_axis);
auto gather_op = ConstructNgNode<Gather>(node.get_name(), ng_input, ng_input_coords, ng_axis);
return {gather_op};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,23 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateIdentityOp(const NodeContext& node) {
return {node.get_ng_input(0)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,27 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateLogSoftmaxOp(const NodeContext& node) {
auto ng_inp = node.get_ng_input(0);
auto inp_shape = ng_inp.get_shape();
size_t rank = inp_shape.size();
int64_t axes = rank - 1;
return {ConstructNgNode<LogSoftmax>(node.get_name(), ng_inp, axes)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,71 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset7.hpp>
using namespace std;
using namespace ov;
using namespace ov::frontend::tf::detail;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateMaxPoolOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0);
auto tf_strides = node.get_attribute<std::vector<int32_t>>("strides");
auto tf_ksize = node.get_attribute<std::vector<int32_t>>("ksize");
auto tf_padding_type = node.get_attribute<std::string>("padding");
auto tf_data_format = node.get_attribute<std::string>("data_format");
bool is_nhwc = (tf_data_format == "NHWC") || (tf_data_format == "NDHWC");
int N = 2;
if (node.get_name() == "MaxPool3D") {
N = 3;
}
Strides ng_strides(N);
Shape ng_image_shape(N);
Shape ng_kernel_shape(N);
Shape ng_dilations(N, 1);
NHWCtoHW(is_nhwc, tf_strides, ng_strides);
NHWCtoHW(is_nhwc, ng_input.get_shape(), ng_image_shape);
NHWCtoHW(is_nhwc, tf_ksize, ng_kernel_shape);
NHWCtoNCHW(node.get_name(), is_nhwc, ng_input);
CoordinateDiff padding_below;
CoordinateDiff padding_above;
MakePadding(tf_padding_type,
ng_image_shape,
ng_kernel_shape,
ng_strides,
ng_dilations,
padding_below,
padding_above);
// TODO: remove this once nGraph supports negative padding
// (CoordinateDiff) for MaxPool
Shape ng_padding_below(padding_below.begin(), padding_below.end());
Shape ng_padding_above(padding_above.begin(), padding_above.end());
auto ng_maxpool = ConstructNgNode<ov::opset7::MaxPool>(node.get_name(),
ng_input,
ng_strides,
ng_padding_below,
ng_padding_above,
ng_kernel_shape,
ov::op::RoundingType::FLOOR);
NCHWtoNHWC(node.get_name(), is_nhwc, ng_maxpool);
return {ng_maxpool};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,28 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector NoOp(const NodeContext& node) {
if (node.get_ng_input_size() == 0) {
return OutputVector{};
}
if (node.get_ng_input_size() != 1) {
throw errors::InvalidArgument("NoOp has " + to_string(node.get_ng_input_size()) + " inputs, should have 1");
}
return OutputVector{node.get_ng_input(0)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,30 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateOneHotOp(const NodeContext& node) {
auto ng_features = node.get_ng_input(0);
auto ng_depth = node.get_ng_input(1);
auto ng_on = node.get_ng_input(2);
auto ng_off = node.get_ng_input(3);
auto one_hot_axis = node.get_attribute<int64_t>("axis");
auto ng_onehot = make_shared<OneHot>(ng_features, ng_depth, ng_on, ng_off, one_hot_axis);
ng_onehot->set_friendly_name(ng_onehot->get_friendly_name());
return ng_onehot->outputs();
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,73 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
// 3 different Pad Ops: Pad, PadV2, MirrorPad
// See https://www.tensorflow.org/api_docs/cc/class/tensorflow/ops/pad
// See https://www.tensorflow.org/api_docs/cc/class/tensorflow/ops/pad-v2
// See https://www.tensorflow.org/api_docs/cc/class/tensorflow/ops/mirror-pad
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslatePadOp(const NodeContext& node) {
auto ng_input = node.get_ng_input(0), ng_paddings_op = node.get_ng_input(1);
Output<Node> pad_val_op;
// Set inputs and pad_val_op
auto op_type = node.get_op_type();
if (op_type == "Pad" || op_type == "MirrorPad") {
pad_val_op =
ConstructNgNode<Constant>(node.get_name(), ng_input.get_element_type(), Shape(), std::vector<int>({0}));
} else if (op_type == "PadV2") {
pad_val_op = node.get_ng_input(2);
} else {
throw errors::InvalidArgument("Incorrect TF Pad OpType: " + node.get_op_type());
}
// Set pad_mode
auto pad_mode = ov::op::PadMode::CONSTANT;
if (op_type == "MirrorPad") {
auto pad_mode_str = node.get_attribute<std::string>("mode");
if (pad_mode_str == "REFLECT") {
pad_mode = ov::op::PadMode::REFLECT;
} else if (pad_mode_str == "SYMMETRIC") {
pad_mode = ov::op::PadMode::SYMMETRIC;
} else {
throw errors::InvalidArgument(pad_mode_str + " is not an allowed padding mode.");
}
}
// Set pads_begin & pads_end (from the pad_val_op)
std::vector<int64_t> paddings;
GetStaticInputVector(node, 1, &paddings);
if (paddings.size() % 2 != 0) {
throw errors::InvalidArgument("Constant node for paddings does not have an even number of "
"elements");
}
std::vector<int64_t> pad_begin(paddings.size() / 2);
std::vector<int64_t> pad_end(paddings.size() / 2);
for (size_t i = 0; i < paddings.size() / 2; i++) {
pad_begin[i] = paddings[2 * i];
pad_end[i] = paddings[2 * i + 1];
}
auto pads_begin_node = ConstructNgNode<Constant>(node.get_name(), element::i64, Shape{pad_begin.size()}, pad_begin);
auto pads_end_node = ConstructNgNode<Constant>(node.get_name(), element::i64, Shape{pad_end.size()}, pad_end);
// Create final Op
auto result_pad_op =
ConstructNgNode<Pad>(node.get_name(), ng_input, pads_begin_node, pads_end_node, pad_val_op, pad_mode);
return {result_pad_op};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,24 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector PlaceholderOp(const NodeContext& node) {
auto ng_et = node.get_attribute<ov::element::Type>("dtype");
auto ng_shape = node.get_attribute<ov::PartialShape>("shape", ov::PartialShape());
return {ConstructNgNode<Parameter>(node.get_name(), ng_et, ng_shape)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,29 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
ov::OutputVector TranslateRandomUniformOp(const NodeContext& node) {
auto data = node.get_ng_input(0);
auto seed = node.get_attribute<int64_t>("seed");
auto seed2 = node.get_attribute<int64_t>("seed2");
auto minval_const = make_shared<Constant>(element::f32, Shape{}, 0);
auto maxval_const = make_shared<Constant>(element::f32, Shape{}, 1);
auto ng_et = node.get_attribute<ov::element::Type>("dtype");
auto random_uniform = std::make_shared<RandomUniform>(data, minval_const, maxval_const, ng_et, seed, seed2);
random_uniform->set_friendly_name(node.get_name());
return random_uniform->outputs();
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,67 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateReduceOp(const NodeContext& node,
std::function<Output<Node>(Output<Node>, Output<Node>, const bool)> create_ng_node) {
Output<Node> ng_input = node.get_ng_input(0);
auto tf_keep_dims = node.get_attribute<bool>("keep_dims", false);
std::vector<int64_t> axes;
GetStaticInputVector(node, 1, &axes);
Shape input_shape = ng_input.get_shape();
size_t input_rank = input_shape.size();
TF_RETURN_IF_ERROR(CheckAxisDimInRange(axes, input_rank));
std::vector<size_t> ng_reduction_axes_vect(axes.size());
std::transform(axes.begin(), axes.end(), ng_reduction_axes_vect.begin(), [input_rank](int idx) {
return idx + (idx < 0 ? (int)input_rank : 0);
});
auto ng_reduction_axes = ConstructNgNode<Constant>(node.get_name(),
element::i64,
Shape{ng_reduction_axes_vect.size()},
ng_reduction_axes_vect);
Output<Node> ng_node = create_ng_node(ng_input, ng_reduction_axes, tf_keep_dims);
return {ng_node};
}
template <typename T>
OutputVector TranslateDirectReduceOp(const NodeContext& node) {
// ensure its either an arithmetic or a logical reduction
if (!(std::is_base_of<ov::op::util::ArithmeticReduction, T>::value ||
std::is_base_of<ov::op::util::LogicalReduction, T>::value)) {
throw errors::InvalidArgument("Expected node to be either a valid logical or arithmetic reduction "
"type");
}
return TranslateReduceOp(node,
[&node](Output<Node> ng_input, Output<Node> ng_reduction_axes, const bool keep_dims) {
return ConstructNgNode<T>(node.get_name(), ng_input, ng_reduction_axes, keep_dims);
});
}
template OutputVector TranslateDirectReduceOp<ReduceLogicalOr>(const NodeContext& node);
template OutputVector TranslateDirectReduceOp<ReduceLogicalAnd>(const NodeContext& node);
template OutputVector TranslateDirectReduceOp<ReduceMax>(const NodeContext& node);
template OutputVector TranslateDirectReduceOp<ReduceMean>(const NodeContext& node);
template OutputVector TranslateDirectReduceOp<ReduceMin>(const NodeContext& node);
template OutputVector TranslateDirectReduceOp<ReduceProd>(const NodeContext& node);
template OutputVector TranslateDirectReduceOp<ReduceSum>(const NodeContext& node);
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,24 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
ov::OutputVector TranslateRelu6Op(const NodeContext& node) {
auto data = node.get_ng_input(0);
auto clamp = std::make_shared<Clamp>(data, 0.0, 6.0f);
clamp->set_friendly_name(node.get_name());
return clamp->outputs();
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,30 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector RetvalOp(const NodeContext& node) {
// Make sure that this _Retval only has one input node.
if (node.get_ng_input_size() != 1) {
throw errors::InvalidArgument("_Retval has " + to_string(node.get_ng_input_size()) + " inputs, should have 1");
}
// auto ret_val_index = node.get_attribute<int>("index");
// TODO: Put ret_val_index to RT info that should be later utilized to order outpus by indices
return {ConstructNgNode<Result>(node.get_name(), node.get_ng_input(0))};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,28 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov;
using namespace ov::opset8;
using namespace ov::frontend::tf;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
ov::OutputVector TranslateRollOp(const NodeContext& node) {
auto data = node.get_ng_input(0);
auto shift = node.get_ng_input(1);
auto axis = node.get_ng_input(2);
auto roll = std::make_shared<Roll>(data, shift, axis);
roll->set_friendly_name(node.get_name());
return roll->outputs();
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,26 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateRsqrtOp(const NodeContext& node) {
auto input = node.get_ng_input(0);
auto ng_exponent = ConstructNgNode<Constant>(node.get_name(), input.get_element_type(), Shape{1}, -0.5f);
auto power = make_shared<Power>(input, ng_exponent);
power->set_friendly_name(node.get_name());
return power->outputs();
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,39 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateSelectOp(const NodeContext& node) {
TF_OP_VALIDATION_CHECK(node, node.get_all_ng_inputs().size() == 3, "Select op cannot be converted");
auto in_1 = node.get_ng_input(0);
auto in_2 = node.get_ng_input(1);
auto in_3 = node.get_ng_input(2);
if (in_1.get_partial_shape().is_static() && in_2.get_partial_shape().is_static()) {
// select broadcast
if (in_1.get_shape().size() == 1 && in_2.get_shape().size() > 1) {
std::vector<uint64_t> axes(in_2.get_shape().size() - 1);
std::iota(axes.begin(), axes.end(), 1);
auto unsqueeze_axes = make_shared<Constant>(ov::element::i64, Shape{in_2.get_shape().size() - 1}, axes);
auto unsqueeze = make_shared<Unsqueeze>(in_1, unsqueeze_axes);
auto ng_select = make_shared<Select>(unsqueeze, in_2, in_3);
ng_select->set_friendly_name(node.get_name());
return ng_select->outputs();
}
}
auto ng_select = make_shared<Select>(in_1, in_2, in_3);
ng_select->set_friendly_name(node.get_name());
return ng_select->outputs();
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,30 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateSoftmaxOp(const NodeContext& node) {
auto ng_inp = node.get_ng_input(0);
auto inp_shape = ng_inp.get_shape();
size_t rank = inp_shape.size();
int64_t axes = rank - 1;
if (rank < 1) {
throw errors::InvalidArgument("TF Softmax logits must be >=1 dimension");
}
return {ConstructNgNode<Softmax>(node.get_name(), ng_inp, axes)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,26 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateSqrtOp(const NodeContext& node) {
auto input = node.get_ng_input(0);
auto ng_exponent = ConstructNgNode<Constant>(node.get_name(), input.get_element_type(), Shape{1}, 0.5f);
auto power = make_shared<Power>(input, ng_exponent);
power->set_friendly_name(node.get_name());
return power->outputs();
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,35 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
OutputVector TranslateSqueezeOp(const NodeContext& node) {
Output<Node> ng_input = node.get_ng_input(0);
size_t input_dims = ng_input.get_shape().size();
auto tf_axis = node.get_attribute<std::vector<int32_t>>("squeeze_dims");
// If input dimension is negative, make it positive
for (size_t i = 0; i < tf_axis.size(); i++) {
tf_axis[i] = tf_axis[i] < 0 ? (int32_t)(input_dims) + tf_axis[i] : tf_axis[i];
}
auto ng_const = ConstructNgNode<Constant>(node.get_name(), element::i32, Shape{tf_axis.size()}, tf_axis);
return {ConstructNgNode<Squeeze>(node.get_name(), ng_input, ng_const)};
}
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,94 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <op_table.hpp>
#include <openvino/opsets/opset8.hpp>
using namespace std;
using namespace ov::opset8;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
// Helper function to translate a unary op.
//
// Parameters:
//
// TFNodeDecoder* op - TF op being translated. Must have one input.
// const std::vector<const ov::frontend::tf::detail::TensorWrapper*>& static_input_map
// - the static input map
// Builder::OpMap& ng_op_map - The TF-to-nGraph op map.
//
// std::function<Output<Node>(Output<Node>>
// create_unary_op - Function to construct the graph implementing
// the unary op, given the input to the unop
// as an argument.
//
// Example usage:
//
// if (n->type_string == "Square") {
// TF_RETURN_IF_ERROR(TranslateUnaryOp(n, static_input_map, ng_op_map,
// [] (Output<Node> n) {
// return
// (Output<Multiply>(n,n));
// });
// }
OutputVector TranslateUnaryOp(const NodeContext& op, std::function<Output<Node>(Output<Node>)> create_unary_op) {
Output<Node> ng_input = op.get_ng_input(0);
auto ng_node = create_unary_op(ng_input);
if (ng_node != ng_input) {
SetTracingInfo(op.get_name(), ng_node);
}
// SaveNgOp(ng_op_map, node.get_name(), ng_node);
// return Status::OK();
return {ng_node};
}
// Helper function to translate a unary op in cases where there is a one-to-one
// mapping from TensorFlow ops to nGraph ops.
//
// Example usage:
//
// if (n->type_string == "Abs") {
// TF_RETURN_IF_ERROR(TranslateUnaryOp<op::Abs>(n, static_input_map,
// ng_op_map));
// }
//
template <typename T>
OutputVector TranslateUnaryOp(const NodeContext& node) {
return TranslateUnaryOp(node, [&node](Output<Node> n) {
return ConstructNgNode<T>(node.get_name(), n);
});
}
template OutputVector TranslateUnaryOp<Abs>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Acos>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Acosh>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Asin>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Asinh>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Atan>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Atanh>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Ceiling>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Cos>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Cosh>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Exp>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Floor>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Log>(const NodeContext& node);
template OutputVector TranslateUnaryOp<LogicalNot>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Negative>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Relu>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Sigmoid>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Sin>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Sinh>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Sign>(const NodeContext& node);
template OutputVector TranslateUnaryOp<SoftPlus>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Tan>(const NodeContext& node);
template OutputVector TranslateUnaryOp<Tanh>(const NodeContext& node);
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,228 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include "op_table.hpp"
using namespace std;
using namespace ov;
using namespace ov::frontend::tf;
namespace ov {
namespace frontend {
namespace tf {
namespace op {
#define OP_CONVERTER(op) ov::OutputVector op(const NodeContext& node)
#define OP_T_CONVERTER(op) \
template <class T> \
ov::OutputVector op(const NodeContext& node)
OP_T_CONVERTER(TranslateUnaryOp);
OP_T_CONVERTER(TranslateBinaryOp);
OP_T_CONVERTER(TranslateDirectReduceOp);
OP_CONVERTER(TranslateAddNOp);
OP_CONVERTER(TranslateArgMaxOp);
OP_CONVERTER(TranslateArgMinOp);
OP_CONVERTER(TranslateAvgPoolOp);
OP_CONVERTER(TranslateBiasAddOp);
OP_CONVERTER(TranslateCastOp);
OP_CONVERTER(TranslateConcatV2Op);
OP_CONVERTER(TranslateConstOp);
OP_CONVERTER(TranslateConv2DOp);
OP_CONVERTER(TranslateConv2DBackpropInputOp);
OP_CONVERTER(TranslateConv3DOp);
OP_CONVERTER(TranslateCumsumOp);
OP_CONVERTER(TranslateDepthToSpaceOp);
OP_CONVERTER(TranslateDepthwiseConv2dNativeOp);
OP_CONVERTER(TranslateEluOp);
OP_CONVERTER(TranslateExpandDimsOp);
OP_CONVERTER(TranslateFillOp);
OP_CONVERTER(TranslateFloorDivOp);
OP_CONVERTER(TranslateFusedBatchNormOp);
OP_CONVERTER(TranslateGatherOp);
OP_CONVERTER(TranslateGatherV2Op);
OP_CONVERTER(TranslateFusedConv2DOp);
OP_CONVERTER(TranslateFusedMatMulOp);
OP_CONVERTER(TranslateIdentityOp);
// OP_CONVERTER(TranslateIsFiniteOp);
// OP_CONVERTER(TranslateL2LossOp);
OP_CONVERTER(TranslateLogSoftmaxOp);
// OP_CONVERTER(TranslateLog1pOp);
// OP_CONVERTER(TranslateLRNOp);
// OP_CONVERTER(TranslateMatMulOp);
OP_CONVERTER(TranslateMaxPoolOp);
OP_CONVERTER(TranslateNonMaxSuppressionV2Op);
OP_CONVERTER(TranslatePadOp);
OP_CONVERTER(PlaceholderOp);
OP_CONVERTER(NoOp);
OP_CONVERTER(TranslateOneHotOp);
// OP_CONVERTER(TranslatePackOp);
OP_CONVERTER(TranslateRangeOp);
OP_CONVERTER(TranslateRankOp);
OP_CONVERTER(TranslateRandomUniformOp);
OP_CONVERTER(TranslateRelu6Op);
// OP_CONVERTER(TranslateReciprocalOp);
// OP_CONVERTER(TranslateReshapeOp);
OP_CONVERTER(RetvalOp);
OP_CONVERTER(TranslateRollOp);
OP_CONVERTER(TranslateRsqrtOp);
OP_CONVERTER(TranslateSelectOp);
// OP_CONVERTER(TranslateShapeOp);
// OP_CONVERTER(TranslateSizeOp);
// OP_CONVERTER(TranslateSliceOp);
// OP_CONVERTER(transpose2);
OP_CONVERTER(TranslateSoftmaxOp);
// OP_CONVERTER(TranslateSpaceToDepthOp);
// OP_CONVERTER(TranslateSplitOp);
// OP_CONVERTER(TranslateSplitOp);
OP_CONVERTER(TranslateSqueezeOp);
// OP_CONVERTER(TranslateStridedSliceOp);
OP_CONVERTER(TranslateSqrtOp);
// OP_CONVERTER(TranslateTileOp);
// OP_CONVERTER(TranslateTopKV2Op);
// OP_CONVERTER(TranslateTransposeOp);
// OP_CONVERTER(TranslateUnpackOp);
// OP_CONVERTER(TranslateWhereOp);
// OP_CONVERTER(TranslateXdivyOp);
// OP_CONVERTER(TranslateZerosLikeOp);
const std::map<const std::string, const CreatorFunction> get_supported_ops() {
return {
// note: UnaryOp translator declaration for each op must to be added in unary_op.cpp file
{"Abs", TranslateUnaryOp<opset8::Abs>},
{"Acos", TranslateUnaryOp<opset8::Acos>},
{"Acosh", TranslateUnaryOp<opset8::Acosh>},
{"Asin", TranslateUnaryOp<opset8::Asin>},
{"Asinh", TranslateUnaryOp<opset8::Asinh>},
{"Atan", TranslateUnaryOp<opset8::Atan>},
{"Atanh", TranslateUnaryOp<opset8::Atanh>},
{"Ceil", TranslateUnaryOp<opset8::Ceiling>},
{"Cos", TranslateUnaryOp<opset8::Cos>},
{"Cosh", TranslateUnaryOp<opset8::Cosh>},
{"Exp", TranslateUnaryOp<opset8::Exp>},
{"Floor", TranslateUnaryOp<opset8::Floor>},
{"Log", TranslateUnaryOp<opset8::Log>},
{"LogicalNot", TranslateUnaryOp<opset8::LogicalNot>},
{"Neg", TranslateUnaryOp<opset8::Negative>},
{"Relu", TranslateUnaryOp<opset8::Relu>},
{"Sigmoid", TranslateUnaryOp<opset8::Sigmoid>},
{"Sin", TranslateUnaryOp<opset8::Sin>},
{"Sinh", TranslateUnaryOp<opset8::Sinh>},
{"Sign", TranslateUnaryOp<opset8::Sign>},
{"Softplus", TranslateUnaryOp<opset8::SoftPlus>},
{"Tan", TranslateUnaryOp<opset8::Tan>},
{"Tanh", TranslateUnaryOp<opset8::Tanh>},
// note: BinaryOp translator declaration for each op must to be added in binary_op.cpp file
{"Add", TranslateBinaryOp<opset8::Add>},
{"AddV2", TranslateBinaryOp<opset8::Add>},
{"Equal", TranslateBinaryOp<opset8::Equal>},
{"FloorMod", TranslateBinaryOp<opset8::FloorMod>},
{"Greater", TranslateBinaryOp<opset8::Greater>},
{"GreaterEqual", TranslateBinaryOp<opset8::GreaterEqual>},
{"Less", TranslateBinaryOp<opset8::Less>},
{"LessEqual", TranslateBinaryOp<opset8::LessEqual>},
{"LogicalAnd", TranslateBinaryOp<opset8::LogicalAnd>},
{"LogicalOr", TranslateBinaryOp<opset8::LogicalOr>},
{"Maximum", TranslateBinaryOp<opset8::Maximum>},
{"Minimum", TranslateBinaryOp<opset8::Minimum>},
{"Mul", TranslateBinaryOp<opset8::Multiply>},
{"Mod", TranslateBinaryOp<opset8::Mod>},
{"NotEqual", TranslateBinaryOp<opset8::NotEqual>},
{"Pow", TranslateBinaryOp<opset8::Power>},
{"RealDiv", TranslateBinaryOp<opset8::Divide>},
{"SquaredDifference", TranslateBinaryOp<opset8::SquaredDifference>},
{"Sub", TranslateBinaryOp<opset8::Subtract>},
// note: ReduceOp translator declaration for each op must to be added in reduce.cpp file
{"Any", TranslateDirectReduceOp<opset8::ReduceLogicalOr>},
{"All", TranslateDirectReduceOp<opset8::ReduceLogicalAnd>},
{"Max", TranslateDirectReduceOp<opset8::ReduceMax>},
{"Mean", TranslateDirectReduceOp<opset8::ReduceMean>},
{"Min", TranslateDirectReduceOp<opset8::ReduceMin>},
{"Prod", TranslateDirectReduceOp<opset8::ReduceProd>},
{"Sum", TranslateDirectReduceOp<opset8::ReduceSum>},
// Separate translators:
{"AddN", TranslateAddNOp},
{"ArgMax", TranslateArgMaxOp},
{"ArgMin", TranslateArgMinOp},
{"AvgPool", TranslateAvgPoolOp},
{"BiasAdd", TranslateBiasAddOp},
{"Cast", TranslateCastOp},
{"ConcatV2", TranslateConcatV2Op},
{"Const", TranslateConstOp},
{"Conv2D", TranslateConv2DOp},
{"Conv2DBackpropInput", TranslateConv2DBackpropInputOp},
{"Conv3D", TranslateConv3DOp},
{"Cumsum", TranslateCumsumOp},
{"DepthToSpace", TranslateDepthToSpaceOp},
{"DepthwiseConv2dNative", TranslateDepthwiseConv2dNativeOp},
{"Elu", TranslateEluOp},
{"ExpandDims", TranslateExpandDimsOp},
{"Fill", TranslateFillOp},
{"FloorDiv", TranslateFloorDivOp},
{"FusedBatchNorm", TranslateFusedBatchNormOp},
{"FusedBatchNormV2", TranslateFusedBatchNormOp},
{"FusedBatchNormV3", TranslateFusedBatchNormOp},
{"Gather", TranslateGatherOp},
{"GatherV2", TranslateGatherV2Op},
{"_FusedConv2D", TranslateFusedConv2DOp},
{"_FusedMatMul", TranslateFusedMatMulOp},
{"Identity", TranslateIdentityOp},
//{"IsFinite", TranslateIsFiniteOp},
//{"L2Loss", TranslateL2LossOp},
{"LogSoftmax", TranslateLogSoftmaxOp},
//{"Log1p", TranslateLog1pOp},
//{"LRN", TranslateLRNOp},
//{"MatMul", TranslateMatMulOp},
{"MaxPool", TranslateMaxPoolOp},
{"MaxPool3D", TranslateMaxPoolOp},
//{"NonMaxSuppressionV2", TranslateNonMaxSuppressionV2Op},
{"MirrorPad", TranslatePadOp},
{"NoOp", NoOp}, // do nothing
{"OneHot", TranslateOneHotOp},
//{"Pack", TranslatePackOp},
{"Pad", TranslatePadOp},
{"PadV2", TranslatePadOp},
//{"_Arg", ArgOp}, // should be registered as an extension in OVTF
{"Placeholder", PlaceholderOp},
// PreventGradient is just Identity in dataflow terms, so reuse that.
{"PreventGradient", TranslateIdentityOp},
//{"Range", TranslateRangeOp},
//{"Rank", TranslateRankOp},
{"RandomUniform", TranslateRandomUniformOp},
//{"Reciprocal", TranslateReciprocalOp},
{"Relu6", TranslateRelu6Op},
//{"Reshape", TranslateReshapeOp},
{"_Retval", RetvalOp},
{"Roll", TranslateRollOp},
{"Rsqrt", TranslateRsqrtOp},
{"Select", TranslateSelectOp},
{"SelectV2", TranslateSelectOp},
//{"Shape", TranslateShapeOp},
//{"Size", TranslateSizeOp},
//{"Slice", TranslateSliceOp},
//{"Snapshot", TranslateIdentityOp},
{"Softmax", TranslateSoftmaxOp},
//{"SpaceToDepth", TranslateSpaceToDepthOp},
//{"Split", TranslateSplitOp},
//{"SplitV", TranslateSplitVOp},
{"Sqrt", TranslateSqrtOp},
//{"Square", TranslateSquareOp},
{"Squeeze", TranslateSqueezeOp},
//{"StridedSlice", TranslateStridedSliceOp},
//{"Tile", TranslateTileOp},
//{"TopKV2", TranslateTopKV2Op},
//{"Transpose", TranslateTransposeOp},
//{"Unpack", TranslateUnpackOp},
//{"Where", TranslateWhereOp},
//{"Xdivy", TranslateXdivyOp},
//{"ZerosLike", TranslateZerosLikeOp},
};
};
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,26 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <functional>
#include <map>
#include <openvino/core/node_vector.hpp>
#include <string>
#include "ngraph_conversions.hpp"
#include "node_context.hpp"
#include "utils.hpp"
namespace ov {
namespace frontend {
namespace tf {
namespace op {
using CreatorFunction = std::function<::ov::OutputVector(const ::ov::frontend::tf::NodeContext&)>;
const std::map<const std::string, const CreatorFunction> get_supported_ops();
} // namespace op
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,308 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <frontend_manager/frontend_exceptions.hpp>
#include <tensorflow_frontend/place.hpp>
#include "node_context.hpp"
#include "op_def.pb.h"
#include "tensor.pb.h"
#include "types.pb.h"
namespace ov {
namespace frontend {
bool PlaceTF::is_input() const {
const auto& model_ins = m_input_model.get_inputs();
const auto cmp = [this](const ngraph::frontend::Place::Ptr& p) {
return p.get() == this;
};
return std::find_if(model_ins.begin(), model_ins.end(), cmp) != model_ins.end();
}
bool PlaceTF::is_output() const {
const auto& model_outs = m_input_model.get_outputs();
const auto cmp = [this](const ngraph::frontend::Place::Ptr& p) {
return p.get() == this;
};
return std::find_if(model_outs.begin(), model_outs.end(), cmp) != model_outs.end();
}
OpPlaceTF::OpPlaceTF(const ngraph::frontend::InputModel& input_model, std::shared_ptr<DecoderBase> op_decoder)
: PlaceTF(input_model, {op_decoder->get_op_name()}),
m_op_decoder(op_decoder) {}
const std::vector<std::shared_ptr<OutPortPlaceTF>>& OpPlaceTF::get_output_ports() const {
return m_output_ports;
}
const std::map<std::string, std::vector<std::shared_ptr<InPortPlaceTF>>>& OpPlaceTF::get_input_ports() const {
return m_input_ports;
}
std::shared_ptr<InPortPlaceTF> OpPlaceTF::get_input_port_tf(const std::string& inputName, int inputPortIndex) const {
FRONT_END_GENERAL_CHECK(inputPortIndex <= m_input_ports.at(inputName).size(), "inputPortIndex is out of bounds.");
return m_input_ports.at(inputName)[inputPortIndex];
}
std::shared_ptr<DecoderBase> OpPlaceTF::get_decoder() const {
return m_op_decoder;
}
void OpPlaceTF::add_out_port(const std::shared_ptr<OutPortPlaceTF>& output, int idx) {
while (idx >= m_output_ports.size()) {
m_output_ports.push_back(std::shared_ptr<OutPortPlaceTF>());
}
m_output_ports[idx] = output;
}
void OpPlaceTF::add_in_port(const std::shared_ptr<InPortPlaceTF>& input, const std::string& name) {
m_input_ports[name].push_back(input);
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_input_port(const std::string& name) const {
FRONT_END_GENERAL_CHECK(m_input_ports.at(name).size() == 1, "Only one input port should exist.");
return m_input_ports.at(name)[0];
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_input_port(int outputPortIndex) const {
FRONT_END_GENERAL_CHECK(m_input_ports.size() == 1, "Only one named input port should exist.");
return m_input_ports.begin()->second[outputPortIndex];
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_output_port(int outputPortIndex) const {
FRONT_END_GENERAL_CHECK(m_output_ports.size() > outputPortIndex, "No port with index: ", outputPortIndex);
return m_output_ports[outputPortIndex];
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_output_port() const {
FRONT_END_GENERAL_CHECK(m_output_ports.size() == 1, "Only one output port should exist.");
return m_output_ports[0];
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_input_port() const {
FRONT_END_GENERAL_CHECK(m_input_ports.size() == 1 && m_input_ports.begin()->second.size() == 1,
"Only one input port should exist.");
return m_input_ports.begin()->second[0];
}
std::vector<ngraph::frontend::Place::Ptr> OpPlaceTF::get_consuming_operations() const {
std::vector<ngraph::frontend::Place::Ptr> consuming_ops;
for (const auto& out_port : m_output_ports) {
auto consuming_ops_out = out_port->get_consuming_operations();
consuming_ops.insert(consuming_ops.end(), consuming_ops_out.begin(), consuming_ops_out.end());
}
return consuming_ops;
}
std::vector<ngraph::frontend::Place::Ptr> OpPlaceTF::get_consuming_operations(int outputPortIndex) const {
return get_output_port(outputPortIndex)->get_consuming_operations();
}
std::vector<ngraph::frontend::Place::Ptr> OpPlaceTF::get_consuming_ports() const {
std::vector<ngraph::frontend::Place::Ptr> consuming_ports;
for (const auto& out_port : m_output_ports) {
auto consuming_ops_out = out_port->get_consuming_ports();
consuming_ports.insert(consuming_ports.end(), consuming_ops_out.begin(), consuming_ops_out.end());
}
return consuming_ports;
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_input_port(const std::string& inputName, int inputPortIndex) const {
FRONT_END_GENERAL_CHECK(inputPortIndex <= m_input_ports.at(inputName).size(), "inputPortIndex is out of bounds.");
return m_input_ports.at(inputName)[inputPortIndex];
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_source_tensor() const {
return get_input_port()->get_source_tensor();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_source_tensor(const std::string& inputName) const {
return get_input_port(inputName)->get_source_tensor();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_source_tensor(int inputPortIndex) const {
return get_input_port(inputPortIndex)->get_source_tensor();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_source_tensor(const std::string& inputName, int inputPortIndex) const {
return get_input_port(inputName, inputPortIndex)->get_source_tensor();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_target_tensor() const {
return get_output_port()->get_target_tensor();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_producing_operation(const std::string& inputName) const {
return get_input_port(inputName)->get_producing_operation();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_producing_operation(const std::string& inputName,
int inputPortIndex) const {
return get_input_port(inputName, inputPortIndex)->get_producing_operation();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_producing_operation() const {
return get_input_port()->get_producing_operation();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_producing_operation(int inputPortIndex) const {
return get_input_port(inputPortIndex)->get_producing_operation();
}
ngraph::frontend::Place::Ptr OpPlaceTF::get_target_tensor(int outputPortIndex) const {
return get_output_port(outputPortIndex)->get_target_tensor();
}
TensorPlaceTF::TensorPlaceTF(const ngraph::frontend::InputModel& input_model,
const ov::PartialShape& pshape,
ov::element::Type type,
const std::vector<std::string>& names)
: PlaceTF(input_model, names),
m_pshape(pshape),
m_type(type) {}
std::vector<ngraph::frontend::Place::Ptr> TensorPlaceTF::get_consuming_ports() const {
std::vector<ngraph::frontend::Place::Ptr> consuming_ports;
for (const auto& consuming_port : m_consuming_ports) {
if (const auto& locked = consuming_port.lock()) {
consuming_ports.push_back(locked);
} else {
FRONT_END_THROW("Consuming Port has expired.");
}
}
return consuming_ports;
}
ngraph::frontend::Place::Ptr TensorPlaceTF::get_producing_port() const {
FRONT_END_GENERAL_CHECK(m_producing_ports.size() == 1, "Only one producing port is supported.");
if (const auto& producing_port = m_producing_ports[0].lock()) {
return producing_port;
}
FRONT_END_THROW("Producing Port has expired.");
}
void TensorPlaceTF::add_producing_port(const std::shared_ptr<OutPortPlaceTF>& out_port) {
m_producing_ports.push_back(out_port);
}
void TensorPlaceTF::add_consuming_port(const std::shared_ptr<InPortPlaceTF>& in_port) {
m_consuming_ports.push_back(in_port);
}
std::vector<ngraph::frontend::Place::Ptr> TensorPlaceTF::get_consuming_operations() const {
std::vector<ngraph::frontend::Place::Ptr> consuming_ops;
for (const auto& consuming_port : m_consuming_ports) {
if (auto port_ptr = consuming_port.lock()) {
auto port_consuming_ops = port_ptr->get_consuming_operations();
consuming_ops.insert(consuming_ops.end(), port_consuming_ops.begin(), port_consuming_ops.end());
} else {
FRONT_END_THROW("Port has expired.");
}
}
return consuming_ops;
}
bool TensorPlaceTF::is_equal_data(ngraph::frontend::Place::Ptr another) const {
auto consuming_ports = get_consuming_ports();
bool eq_to_consuming_port =
std::any_of(consuming_ports.begin(), consuming_ports.end(), [&another](const Ptr& place) {
return place->is_equal(another);
});
return is_equal(another) || get_producing_port()->is_equal(another) || eq_to_consuming_port;
}
ngraph::frontend::Place::Ptr TensorPlaceTF::get_producing_operation() const {
return get_producing_port()->get_producing_operation();
}
std::shared_ptr<TensorPlaceTF> InPortPlaceTF::get_source_tensor_tf() const {
if (const auto& tensor = m_source_tensor.lock()) {
return tensor;
}
FRONT_END_THROW("Source Tensor has expired.");
}
std::shared_ptr<OpPlaceTF> InPortPlaceTF::get_op() {
if (const auto& op = m_op.lock()) {
return op;
}
FRONT_END_THROW("Operation has expired.");
}
void InPortPlaceTF::set_source_tensor(const std::weak_ptr<TensorPlaceTF>& source_tensor) {
m_source_tensor = source_tensor;
}
std::vector<ngraph::frontend::Place::Ptr> InPortPlaceTF::get_consuming_operations() const {
if (const auto& consuming_op = m_op.lock()) {
return {consuming_op};
}
FRONT_END_THROW("Operation has expired.");
}
ngraph::frontend::Place::Ptr InPortPlaceTF::get_source_tensor() const {
if (const auto& tensor = m_source_tensor.lock()) {
return tensor;
}
FRONT_END_THROW("Source Tensor has expired.");
}
ngraph::frontend::Place::Ptr InPortPlaceTF::get_producing_port() const {
return get_source_tensor()->get_producing_port();
}
bool InPortPlaceTF::is_equal_data(ngraph::frontend::Place::Ptr another) const {
return get_source_tensor()->is_equal_data(another);
}
ngraph::frontend::Place::Ptr InPortPlaceTF::get_producing_operation() const {
return get_producing_port()->get_producing_operation();
}
std::shared_ptr<TensorPlaceTF> OutPortPlaceTF::get_target_tensor_tf() const {
if (const auto& target_tensor = m_target_tensor.lock()) {
return target_tensor;
}
FRONT_END_THROW("Target Tensor has expired.");
}
std::vector<ngraph::frontend::Place::Ptr> OutPortPlaceTF::get_consuming_operations() const {
if (auto tensor_ptr = m_target_tensor.lock()) {
return tensor_ptr->get_consuming_operations();
}
FRONT_END_THROW("Tensor has expired.");
}
void OutPortPlaceTF::set_target_tensor(const std::weak_ptr<TensorPlaceTF>& target_tensor) {
m_target_tensor = target_tensor;
}
std::vector<ngraph::frontend::Place::Ptr> OutPortPlaceTF::get_consuming_ports() const {
if (auto tensor_ptr = m_target_tensor.lock()) {
return tensor_ptr->get_consuming_ports();
}
FRONT_END_THROW("Tensor has expired.");
}
bool OutPortPlaceTF::is_equal_data(ngraph::frontend::Place::Ptr another) const {
return get_target_tensor()->is_equal_data(another);
}
ngraph::frontend::Place::Ptr OutPortPlaceTF::get_target_tensor() const {
if (const auto& target_tensor = m_target_tensor.lock()) {
return target_tensor;
}
FRONT_END_THROW("Target Tensor has expired.");
}
ngraph::frontend::Place::Ptr OutPortPlaceTF::get_producing_operation() const {
if (auto op = m_op.lock()) {
return op;
}
FRONT_END_THROW("Operation has expired.");
}
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,41 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "AllocationDescriptionProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/allocation_description_go_proto";
message AllocationDescription {
// Total number of bytes requested
int64 requested_bytes = 1;
// Total number of bytes allocated if known
int64 allocated_bytes = 2;
// Name of the allocator used
string allocator_name = 3;
// Identifier of the allocated buffer if known
int64 allocation_id = 4;
// Set if this tensor only has one remaining reference
bool has_single_reference = 5;
// Address of the allocation.
uint64 ptr = 6;
}

View File

@ -0,0 +1,148 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
// Defines the text format for including per-op API definition and
// overrides for client language op code generators.
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "ApiDefProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/api_def_go_proto";
import "attr_value.proto";
// Used to specify and override the default API & behavior in the
// generated code for client languages, from what you would get from
// the OpDef alone. There will be a set of ApiDefs that are common
// to all client languages, and another set per client language.
// The per-client-language ApiDefs will inherit values from the
// common ApiDefs which it can either replace or modify.
//
// We separate the API definition from the OpDef so we can evolve the
// API while remaining backwards compatible when interpretting old
// graphs. Overrides go in an "api_def.pbtxt" file with a text-format
// ApiDefs message.
//
// WARNING: Be *very* careful changing the API for any existing op --
// you can change the semantics of existing code. These changes may
// need to wait until a major release of TensorFlow to avoid breaking
// our compatibility promises.
message ApiDef {
// Name of the op (in the OpDef) to specify the API for.
string graph_op_name = 1;
// If this op is deprecated, set deprecation message to the message
// that should be logged when this op is used.
// The message should indicate alternative op to use, if any.
string deprecation_message = 12;
// Major version when the op will be deleted. For e.g. set this
// value to 2 if op API should be removed in TensorFlow 2.0 and
// deprecated in versions before that.
int32 deprecation_version = 13;
enum Visibility {
// Normally this is "VISIBLE" unless you are inheriting a
// different value from another ApiDef.
DEFAULT_VISIBILITY = 0;
// Publicly visible in the API.
VISIBLE = 1;
// Do not include this op in the generated API. If visibility is
// set to 'SKIP', other fields are ignored for this op.
SKIP = 2;
// Hide this op by putting it into an internal namespace (or whatever
// is appropriate in the target language).
HIDDEN = 3;
}
Visibility visibility = 2;
// If you specify any endpoint, this will replace all of the
// inherited endpoints. The first endpoint should be the
// "canonical" endpoint, and should not be deprecated (unless all
// endpoints are deprecated).
message Endpoint {
// Name should be either like "CamelCaseName" or
// "Package.CamelCaseName". Client-language-specific ApiDefs may
// use a snake_case convention instead of CamelCase.
string name = 1;
// Set if this endpoint is deprecated. If set to true, a message suggesting
// to use a non-deprecated endpoint instead will be printed. If all
// endpoints are deprecated, set deprecation_message in ApiDef instead.
bool deprecated = 3;
// Major version when an endpoint will be deleted. For e.g. set this
// value to 2 if endpoint should be removed in TensorFlow 2.0 and
// deprecated in versions before that.
int32 deprecation_version = 4;
}
repeated Endpoint endpoint = 3;
message Arg {
string name = 1;
// Change the name used to access this arg in the API from what
// is used in the GraphDef. Note that these names in `backticks`
// will also be replaced in the summary & description fields.
string rename_to = 2;
// Note: this will replace any inherited arg doc. There is no
// current way of modifying arg descriptions (other than replacing
// them entirely) as can be done with op descriptions.
string description = 3;
}
repeated Arg in_arg = 4;
repeated Arg out_arg = 5;
// List of original in_arg names to specify new argument order.
// Length of arg_order should be either empty to keep current order
// or match size of in_arg.
repeated string arg_order = 11;
// Description of the graph-construction-time configuration of this
// Op. That is to say, this describes the attr fields that will
// be specified in the NodeDef.
message Attr {
string name = 1;
// Change the name used to access this attr in the API from what
// is used in the GraphDef. Note that these names in `backticks`
// will also be replaced in the summary & description fields.
string rename_to = 2;
// Specify a new default value to use for this attr. This default
// will be used when creating new graphs, as opposed to the
// default in the OpDef, which will be used when interpreting old
// GraphDefs.
AttrValue default_value = 3;
// Note: this will replace any inherited attr doc, there is no current
// way of modifying attr descriptions as can be done with op descriptions.
string description = 4;
}
repeated Attr attr = 6;
// One-line human-readable description of what the Op does.
string summary = 7;
// Additional, longer human-readable description of what the Op does.
string description = 8;
// Modify an existing/inherited description by adding text to the beginning
// or end.
string description_prefix = 9;
string description_suffix = 10;
}
message ApiDefs {
repeated ApiDef op = 1;
}

View File

@ -0,0 +1,76 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "tensor.proto";
import "tensor_shape.proto";
import "types.proto";
option cc_enable_arenas = true;
option java_outer_classname = "AttrValueProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/attr_value_go_proto";
// Protocol buffer representing the value for an attr used to configure an Op.
// Comment indicates the corresponding attr type. Only the field matching the
// attr type may be filled.
message AttrValue {
// LINT.IfChange
message ListValue {
repeated bytes s = 2; // "list(string)"
repeated int64 i = 3 [packed = true]; // "list(int)"
repeated float f = 4 [packed = true]; // "list(float)"
repeated bool b = 5 [packed = true]; // "list(bool)"
repeated DataType type = 6 [packed = true]; // "list(type)"
repeated TensorShapeProto shape = 7; // "list(shape)"
repeated TensorProto tensor = 8; // "list(tensor)"
repeated NameAttrList func = 9; // "list(attr)"
}
// LINT.ThenChange(https://www.tensorflow.org/code/tensorflow/c/c_api.cc)
oneof value {
bytes s = 2; // "string"
int64 i = 3; // "int"
float f = 4; // "float"
bool b = 5; // "bool"
DataType type = 6; // "type"
TensorShapeProto shape = 7; // "shape"
TensorProto tensor = 8; // "tensor"
ListValue list = 1; // any "list(...)"
// "func" represents a function. func.name is a function's name or
// a primitive op's name. func.attr.first is the name of an attr
// defined for that function. func.attr.second is the value for
// that attr in the instantiation.
NameAttrList func = 10;
// This is a placeholder only used in nodes defined inside a
// function. It indicates the attr value will be supplied when
// the function is instantiated. For example, let us suppose a
// node "N" in function "FN". "N" has an attr "A" with value
// placeholder = "foo". When FN is instantiated with attr "foo"
// set to "bar", the instantiated node N's attr A will have been
// given the value "bar".
string placeholder = 9;
}
}
// A list of attr names and their values. The whole list is attached
// with a string name. E.g., MatMul[T=float].
message NameAttrList {
string name = 1;
map<string, AttrValue> attr = 2;
}

View File

@ -0,0 +1,101 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "tensor_shape.proto";
import "types.proto";
option cc_enable_arenas = true;
option java_outer_classname = "CostGraphProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/cost_graph_go_proto";
message CostGraphDef {
message Node {
// The name of the node. Names are globally unique.
string name = 1;
// The device of the node. Can be empty if the node is mapped to the
// default partition or partitioning hasn't been run yet.
string device = 2;
// The id of the node. Node ids are only unique inside a partition.
int32 id = 3;
// Inputs of this node. They must be executed before this node can be
// executed. An input is a particular output of another node, specified
// by the node id and the output index.
message InputInfo {
int32 preceding_node = 1;
int32 preceding_port = 2;
}
repeated InputInfo input_info = 4;
// Outputs of this node.
message OutputInfo {
int64 size = 1;
// If >= 0, the output is an alias of an input. Note that an alias input
// may itself be an alias. The algorithm will therefore need to follow
// those pointers.
int64 alias_input_port = 2;
TensorShapeProto shape = 3;
DataType dtype = 4;
}
repeated OutputInfo output_info = 5;
// Temporary memory used by this node.
int64 temporary_memory_size = 6;
// Persistent memory used by this node.
int64 persistent_memory_size = 12;
int64 host_temp_memory_size = 10 [deprecated = true];
int64 device_temp_memory_size = 11 [deprecated = true];
int64 device_persistent_memory_size = 16 [deprecated = true];
// Estimate of the computational cost of this node, in microseconds.
int64 compute_cost = 9;
// Analytical estimate of the computational cost of this node, in
// microseconds.
int64 compute_time = 14;
// Analytical estimate of the memory access cost of this node, in
// microseconds.
int64 memory_time = 15;
// If true, the output is permanent: it can't be discarded, because this
// node is part of the "final output". Nodes may depend on final nodes.
bool is_final = 7;
// Ids of the control inputs for this node.
repeated int32 control_input = 8;
// Are the costs inaccurate?
bool inaccurate = 17;
}
repeated Node node = 1;
// Total cost of this graph, typically used for balancing decisions.
message AggregatedCost {
// Aggregated cost value.
float cost = 1;
// Aggregated cost dimension (e.g. 'memory', 'compute', 'network').
string dimension = 2;
}
repeated AggregatedCost cost = 2;
}

View File

@ -0,0 +1,191 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow.data;
// Represents the type of auto-sharding we enable.
enum AutoShardPolicy {
AUTO = 0;
FILE = 1;
DATA = 2;
OFF = -1;
}
message DistributeOptions {
// The type of sharding that auto-shard should attempt. If this is set to
// FILE, then we will attempt to shard by files (each worker will get a set of
// files to process). If we cannot find a set of files to shard for at least
// one file per worker, we will error out. When this option is selected, make
// sure that you have enough files so that each worker gets at least one file.
// There will be a runtime error thrown if there are insufficient files. If
// this is set to DATA, then we will shard by elements produced by the
// dataset, and each worker will process the whole dataset and discard the
// portion that is not for itself. If this is set to OFF, then we will not
// autoshard, and each worker will receive a copy of the full dataset. This
// option is set to AUTO by default, AUTO will attempt to first shard by FILE,
// and fall back to sharding by DATA if we cannot find a set of files to
// shard.
AutoShardPolicy auto_shard_policy = 1;
// The number of devices attached to this input pipeline.
oneof optional_num_devices {
int32 num_devices = 2;
}
}
message MapVectorization {
// Whether to vectorize map transformations.
oneof optional_enabled {
bool enabled = 1;
}
// Whether to use ChooseFastestBranchDataset with this transformation. If
// True, the pipeline picks between the vectorized and original segment at
// runtime based on their iterations speed.
oneof optional_use_choose_fastest {
bool use_choose_fastest = 2;
}
}
message OptimizationOptions {
// Whether to apply default graph optimizations. If False, only graph
// optimizations that have been explicitly enabled will be applied.
oneof optional_apply_default_optimizations {
bool apply_default_optimizations = 1;
}
// Whether to automatically tune performance knobs.
oneof optional_autotune {
bool autotune = 2;
}
// When autotuning is enabled (through autotune), determines whether to also
// autotune buffer sizes for datasets with parallelism.
oneof optional_autotune_buffers {
bool autotune_buffers = 3;
}
// When autotuning is enabled (through autotune), determines the CPU budget to
// use. Values greater than the number of schedulable CPU cores are allowed
// but may result in CPU contention.
oneof optional_autotune_cpu_budget {
int32 autotune_cpu_budget = 4;
}
// When autotuning is enabled (through autotune), determines the RAM budget to
// use. Values greater than the available RAM in bytes may result in OOM. If
// 0, defaults to half of the available RAM in bytes.
oneof optional_autotune_ram_budget {
int32 autotune_ram_budget = 5;
}
// Whether to fuse filter transformations.
oneof optional_filter_fusion {
bool filter_fusion = 6;
}
// Whether to fuse filter dataset that predicts random_uniform < rate into a
// sampling dataset.
oneof optional_filter_with_random_uniform_fusion {
bool filter_with_random_uniform_fusion = 7;
}
// Whether to hoist tf.random_uniform() ops out of map transformations.
oneof optional_hoist_random_uniform {
bool hoist_random_uniform = 8;
}
// Whether to fuse map and batch transformations.
oneof optional_map_and_batch_fusion {
bool map_and_batch_fusion = 9;
}
// Whether to fuse map and filter transformations.
oneof optional_map_and_filter_fusion {
bool map_and_filter_fusion = 10;
}
// Whether to fuse map transformations.
oneof optional_map_fusion {
bool map_fusion = 11;
}
// Whether to parallelize stateless map transformations.
oneof optional_map_parallelization {
bool map_parallelization = 12;
}
// The map vectorization options associated with the dataset.
MapVectorization map_vectorization = 13;
// Whether to eliminate no-op transformations.
oneof optional_noop_elimination {
bool noop_elimination = 14;
}
// Whether to parallelize copying of batch elements. This optimization is
// highly experimental and can cause performance degradation (e.g. when the
// parallelization overhead exceeds the benefits of performing the data copies
// in parallel). You should only enable this optimization if a) your input
// pipeline is bottlenecked on batching and b) you have validated that this
// optimization improves performance.
oneof optional_parallel_batch {
bool parallel_batch = 15;
}
// Whether to reorder ops that will discard data to the front of unary
// cardinality preserving transformations, e.g. dataset.map(...).take(3) will
// be optimized to dataset.take(3).map(...). For now this optimization will
// move `skip`, `shard` and `take` to the front of `map` and `prefetch`. This
// optimization is only for performance; it will not affect the output of the
// dataset.
oneof optional_reorder_data_discarding_ops {
bool reorder_data_discarding_ops = 16;
}
// Whether to fuse shuffle and repeat transformations.
oneof optional_shuffle_and_repeat_fusion {
bool shuffle_and_repeat_fusion = 17;
}
}
message ThreadingOptions {
// If set, it overrides the maximum degree of intra-op parallelism.
oneof optional_max_intra_op_parallelism {
int32 max_intra_op_parallelism = 1;
}
// If set, the dataset will use a private threadpool of the given size.
oneof optional_private_threadpool_size {
int32 private_threadpool_size = 2;
}
}
// Represents how to handle external state during serialization.
enum ExternalStatePolicy {
WARN = 0;
IGNORE = 1;
FAIL = 2;
}
// Message stored with Dataset objects to control how datasets are processed and
// optimized.
message Options {
// Whether the outputs need to be produced in deterministic order.
oneof optional_deterministic {
bool deterministic = 1;
}
// The distribution strategy options associated with the dataset.
DistributeOptions distribute_options = 2;
// The optimization options associated with the dataset.
OptimizationOptions optimization_options = 3;
// Whether to introduce 'slack' in the last `prefetch` of the input pipeline,
// if it exists. This may reduce CPU contention with accelerator host-side
// activity at the start of a step. The slack frequency is determined by the
// number of devices attached to this input pipeline.
oneof optional_slack {
bool slack = 4;
}
// The threading options associated with the dataset.
ThreadingOptions threading_options = 5;
// This option can be used to override the default policy for how to handle
// external state when serializing a dataset or checkpointing its iterator.
// There are three settings available - IGNORE: External state is ignored
// without a warning; WARN: External state is ignored and a warning is logged;
// FAIL: External state results in an error.
oneof optional_external_state_policy {
ExternalStatePolicy external_state_policy = 6;
}
}

View File

@ -0,0 +1,65 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "DeviceAttributesProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/device_attributes_go_proto";
message InterconnectLink {
int32 device_id = 1;
string type = 2;
int32 strength = 3;
}
message LocalLinks {
repeated InterconnectLink link = 1;
}
message DeviceLocality {
// Optional bus locality of device. Default value of 0 means
// no specific locality. Specific localities are indexed from 1.
int32 bus_id = 1;
// Optional NUMA locality of device.
int32 numa_node = 2;
// Optional local interconnect links to other devices.
LocalLinks links = 3;
}
message DeviceAttributes {
// Fully specified name of the device within a cluster.
string name = 1;
// String representation of device_type.
string device_type = 2;
// Memory capacity of device in bytes.
int64 memory_limit = 4;
// Platform-specific data about device that may be useful
// for supporting efficient data transfers.
DeviceLocality locality = 5;
// A device is assigned a global unique number each time it is
// initialized. "incarnation" should never be 0.
fixed64 incarnation = 6;
// String representation of the physical device that this device maps to.
string physical_device_desc = 7;
}

View File

@ -0,0 +1,138 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "attr_value.proto";
import "node_def.proto";
import "op_def.proto";
option cc_enable_arenas = true;
option java_outer_classname = "FunctionProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/function_go_proto";
// A library is a set of named functions.
message FunctionDefLibrary {
repeated FunctionDef function = 1;
repeated GradientDef gradient = 2;
}
// A function can be instantiated when the runtime can bind every attr
// with a value. When a GraphDef has a call to a function, it must
// have binding for every attr defined in the signature.
//
// TODO(zhifengc):
// * device spec, etc.
message FunctionDef {
// The definition of the function's name, arguments, return values,
// attrs etc.
OpDef signature = 1;
// Attributes specific to this function definition.
map<string, AttrValue> attr = 5;
// Attributes for function arguments. These attributes are the same set of
// valid attributes as to _Arg nodes.
message ArgAttrs {
map<string, AttrValue> attr = 1;
}
map<uint32, ArgAttrs> arg_attr = 7;
// Unique IDs for each resource argument, used to track aliasing resources. If
// Argument A and Argument B alias each other, then
// resource_arg_unique_ids[A.index] == resource_arg_unique_ids[B.index].
//
// If this field is empty, none of the arguments could alias; otherwise, every
// resource argument should have an entry in this field.
//
// When instantiated, the unique IDs will be attached to the _Arg nodes'
// "_resource_arg_unique_id" attribute.
map<uint32, uint32> resource_arg_unique_id = 8;
// NOTE: field id 2 deleted on Jan 11, 2017, GraphDef version 21.
reserved 2;
// In both of the following fields, there is the need to specify an
// output that is used as either the input to another node (in
// `node_def`) or as a return value of the function (in `ret`).
// Unlike the NodeDefs in GraphDef, we need to be able to specify a
// list in some cases (instead of just single outputs). Also, we
// need to be able to deal with lists of unknown length (so the
// output index may not be known at function definition time). So
// we use the following format instead:
// * "fun_in" where "fun_in" is the name of a function input arg in
// the `signature` field above. This represents that input, whether
// it is a single tensor or a list.
// * "fun_in:0" gives the first element of a function input arg (a
// non-list input is considered a list of length 1 for these
// purposes).
// * "node:out" where "node" is the name of a node in `node_def` and
// "out" is the name one of its op's output arguments (the name
// comes from the OpDef of the node's op). This represents that
// node's output, whether it is a single tensor or a list.
// Note: We enforce that an op's output arguments are never
// renamed in the backwards-compatibility test.
// * "node:out:0" gives the first element of a node output arg (a
// non-list output is considered a list of length 1 for these
// purposes).
//
// NOT CURRENTLY SUPPORTED (but may be in the future):
// * "node:out:-1" gives last element in a node output list
// * "node:out:1:" gives a list with all but the first element in a
// node output list
// * "node:out::-1" gives a list with all but the last element in a
// node output list
// The body of the function. Unlike the NodeDefs in a GraphDef, attrs
// may have values of type `placeholder` and the `input` field uses
// the "output" format above.
// By convention, "op" in node_def is resolved by consulting with a
// user-defined library first. If not resolved, "func" is assumed to
// be a builtin op.
repeated NodeDef node_def = 3;
// A mapping from the output arg names from `signature` to the
// outputs from `node_def` that should be returned by the function.
map<string, string> ret = 4;
// A mapping from control output names from `signature` to node names in
// `node_def` which should be control outputs of this function.
map<string, string> control_ret = 6;
}
// GradientDef defines the gradient function of a function defined in
// a function library.
//
// A gradient function g (specified by gradient_func) for a function f
// (specified by function_name) must follow the following:
//
// The function 'f' must be a numerical function which takes N inputs
// and produces M outputs. Its gradient function 'g', which is a
// function taking N + M inputs and produces N outputs.
//
// I.e. if we have
// (y1, y2, ..., y_M) = f(x1, x2, ..., x_N),
// then, g is
// (dL/dx1, dL/dx2, ..., dL/dx_N) = g(x1, x2, ..., x_N,
// dL/dy1, dL/dy2, ..., dL/dy_M),
// where L is a scalar-value function of (x1, x2, ..., xN) (e.g., the
// loss function). dL/dx_i is the partial derivative of L with respect
// to x_i.
message GradientDef {
string function_name = 1; // The function name.
string gradient_func = 2; // The gradient function's name.
}

View File

@ -0,0 +1,68 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "function.proto";
import "node_def.proto";
import "versions.proto";
option cc_enable_arenas = true;
option java_outer_classname = "GraphProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/graph_go_proto";
// Represents the graph of operations
message GraphDef {
repeated NodeDef node = 1;
// Compatibility versions of the graph. See core/public/version.h for version
// history. The GraphDef version is distinct from the TensorFlow version, and
// each release of TensorFlow will support a range of GraphDef versions.
VersionDef versions = 4;
// Deprecated single version field; use versions above instead. Since all
// GraphDef changes before "versions" was introduced were forward
// compatible, this field is entirely ignored.
int32 version = 3 [deprecated = true];
// "library" provides user-defined functions.
//
// Naming:
// * library.function.name are in a flat namespace.
// NOTE: We may need to change it to be hierarchical to support
// different orgs. E.g.,
// { "/google/nn", { ... }},
// { "/google/vision", { ... }}
// { "/org_foo/module_bar", { ... }}
// map<string, FunctionDefLib> named_lib;
// * If node[i].op is the name of one function in "library",
// node[i] is deemed as a function call. Otherwise, node[i].op
// must be a primitive operation supported by the runtime.
//
//
// Function call semantics:
//
// * The callee may start execution as soon as some of its inputs
// are ready. The caller may want to use Tuple() mechanism to
// ensure all inputs are ready in the same time.
//
// * The consumer of return values may start executing as soon as
// the return values the consumer depends on are ready. The
// consumer may want to use Tuple() mechanism to ensure the
// consumer does not start until all return values of the callee
// function are ready.
FunctionDefLibrary library = 2;
}

View File

@ -0,0 +1,83 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "types.proto";
option cc_enable_arenas = true;
option java_outer_classname = "GraphTransferInfoProto";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/graph_transfer_info_go_proto";
message GraphTransferNodeInput {
int32 node_id = 1;
int32 output_port = 2;
}
message GraphTransferNodeInfo {
string name = 1;
int32 node_id = 2;
string type_name = 3;
int32 soc_op_id = 4;
int32 padding_id = 5;
int32 input_count = 6;
int32 output_count = 7;
}
message GraphTransferConstNodeInfo {
string name = 1;
int32 node_id = 2;
repeated int64 shape = 3;
bytes data = 4;
DataType dtype = 5;
}
message GraphTransferNodeInputInfo {
int32 node_id = 1;
repeated GraphTransferNodeInput node_input = 2;
}
message GraphTransferNodeOutputInfo {
int32 node_id = 1;
repeated int32 max_byte_size = 2;
}
message GraphTransferGraphInputNodeInfo {
string name = 1;
repeated int64 shape = 2;
DataType dtype = 3;
}
message GraphTransferGraphOutputNodeInfo {
string name = 1;
repeated int64 shape = 2;
DataType dtype = 3;
}
// Protocol buffer representing a handle to a tensorflow resource. Handles are
// not valid across executions, but can be serialized back and forth from within
// a single run.
message GraphTransferInfo {
enum Destination {
NOP = 0;
HEXAGON = 1;
}
repeated GraphTransferNodeInfo node_info = 1;
repeated GraphTransferConstNodeInfo const_node_info = 2;
repeated GraphTransferNodeInputInfo node_input_info = 3;
repeated GraphTransferNodeOutputInfo node_output_info = 4;
// Input Node parameters of transferred graph
repeated GraphTransferGraphInputNodeInfo graph_input_node_info = 5;
repeated GraphTransferGraphOutputNodeInfo graph_output_node_info = 6;
// Destination of graph transfer
Destination destination = 7;
}

View File

@ -0,0 +1,60 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "attr_value.proto";
option cc_enable_arenas = true;
option java_outer_classname = "KernelDefProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/kernel_def_go_proto";
message KernelDef {
// Must match the name of an Op.
string op = 1;
// Type of device this kernel runs on.
string device_type = 2;
message AttrConstraint {
// Name of an attr from the Op.
string name = 1;
// A list of values that this kernel supports for this attr.
// Like OpDef.AttrDef.allowed_values, except for kernels instead of Ops.
AttrValue allowed_values = 2;
}
repeated AttrConstraint constraint = 3;
// Names of the Op's input_/output_args that reside in host memory
// instead of device memory.
repeated string host_memory_arg = 4;
// This allows experimental kernels to be registered for an op that
// won't be used unless the user specifies a "_kernel" attr with
// value matching this.
string label = 5;
// Prioritization of kernel amongst different devices. By default we assume
// priority is 0. The higher the priority the better. By default (i.e. if
// this is not set), we prefer GPU kernels over CPU.
int32 priority = 6;
}
// A collection of KernelDefs
message KernelList {
repeated KernelDef kernel = 1;
}

View File

@ -0,0 +1,107 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "tensor_description.proto";
option cc_enable_arenas = true;
option java_outer_classname = "LogMemoryProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/log_memory_go_proto";
message MemoryLogStep {
// Process-unique step id.
int64 step_id = 1;
// Handle describing the feeds and fetches of the step.
string handle = 2;
}
message MemoryLogTensorAllocation {
// Process-unique step id.
int64 step_id = 1;
// Name of the kernel making the allocation as set in GraphDef,
// e.g., "affine2/weights/Assign".
string kernel_name = 2;
// Allocated tensor details.
TensorDescription tensor = 3;
}
message MemoryLogTensorDeallocation {
// Id of the tensor buffer being deallocated, used to match to a
// corresponding allocation.
int64 allocation_id = 1;
// Name of the allocator used.
string allocator_name = 2;
}
message MemoryLogTensorOutput {
// Process-unique step id.
int64 step_id = 1;
// Name of the kernel producing an output as set in GraphDef, e.g.,
// "affine2/weights/Assign".
string kernel_name = 2;
// Index of the output being set.
int32 index = 3;
// Output tensor details.
TensorDescription tensor = 4;
}
message MemoryLogRawAllocation {
// Process-unique step id.
int64 step_id = 1;
// Name of the operation making the allocation.
string operation = 2;
// Number of bytes in the allocation.
int64 num_bytes = 3;
// Address of the allocation.
uint64 ptr = 4;
// Id of the tensor buffer being allocated, used to match to a
// corresponding deallocation.
int64 allocation_id = 5;
// Name of the allocator used.
string allocator_name = 6;
}
message MemoryLogRawDeallocation {
// Process-unique step id.
int64 step_id = 1;
// Name of the operation making the deallocation.
string operation = 2;
// Id of the tensor buffer being deallocated, used to match to a
// corresponding allocation.
int64 allocation_id = 3;
// Name of the allocator used.
string allocator_name = 4;
// True if the deallocation is queued and will be performed later,
// e.g. for GPU lazy freeing of buffers.
bool deferred = 5;
}

View File

@ -0,0 +1,142 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow.data.model;
option cc_enable_arenas = true;
// Class of a node in the performance model.
enum NodeClass {
UNKNOWN = 0;
INTERLEAVE_MANY = 1;
ASYNC_INTERLEAVE_MANY = 2;
KNOWN_RATIO = 3;
ASYNC_KNOWN_RATIO = 4;
UNKNOWN_RATIO = 5;
}
// Algorithm used for model autotuning optimization.
enum AutotuneAlgorithm {
HILL_CLIMB = 0;
GRADIENT_DESCENT = 1;
}
// Protocol buffer representing the data used by the autotuning modeling
// framework.
message ModelProto {
// General representation of a node in the model.
message Node {
// Unique node ID.
int64 id = 1;
// Human-readable name of the node.
string name = 2;
// An indication whether autotuning is enabled for this node.
bool autotune = 3;
// The number of bytes stored in this node's buffer.
int64 buffered_bytes = 4;
// The number of elements stored in this node's buffer.
int64 buffered_elements = 5;
// The number of bytes consumed by the node.
int64 bytes_consumed = 6;
// The number of bytes produced by the node.
int64 bytes_produced = 7;
// The number of elements produced by the node.
int64 num_elements = 8;
// The aggregate processing time spent in this node.
int64 processing_time = 9;
// An indication whether this node records metrics about produced and
// consumed elements.
bool record_metrics = 10;
// Represents a node parameter.
message Parameter {
// Human-readable name of the parameter.
string name = 1;
// Identifies the model value of the parameter. This can be different from
// the actual value (e.g. during optimization search).
double value = 2;
// The actual value of the parameter.
double state_value = 3;
// Minimum value of the parameter.
double min = 4;
// Maximum value of the parameter.
double max = 5;
// Identifies whether the parameter should participate in autotuning.
bool tunable = 6;
}
// Parameters of this node.
repeated Parameter parameters = 11;
// Statistic of inputs processing time history.
double input_processing_time_sum = 12;
int64 input_processing_time_count = 13;
// Inputs of this node.
repeated Node inputs = 14;
// Class of this node.
NodeClass node_class = 15;
// Ratio of input to output elements. This is only used by KNOWN_RATIO and
// ASYNC_KNOWN_RATIO nodes.
double ratio = 16;
// Ratio identifies how many parallelism calls are introduced by one
// buffered element. This is only used by ASYNC_KNOWN_RATIO nodes.
double memory_ratio = 17;
}
// Output node of this model.
Node output = 1;
// Counter for node IDs of this model.
int64 id_counter = 2;
// Indicates whether the modeling framework should collect resource usage,
// e.g. CPU, memory.
bool collect_resource_usage = 3;
// Contains parameters of the model autotuning optimization.
message OptimizationParams {
// Algorithm used for autotuning optimization.
AutotuneAlgorithm algorithm = 1;
// Number of available logical threads.
int64 cpu_budget = 2;
// Amount of available memory in bytes.
int64 ram_budget = 3;
// Time between two consecutive `GetNext` calls to the iterator represented
// by the output node.
double model_input_time = 4;
}
OptimizationParams optimization_params = 4;
}

View File

@ -0,0 +1,100 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "attr_value.proto";
option cc_enable_arenas = true;
option java_outer_classname = "NodeProto";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/node_def_go_proto";
message NodeDef {
// The name given to this operator. Used for naming inputs,
// logging, visualization, etc. Unique within a single GraphDef.
// Must match the regexp "[A-Za-z0-9.][A-Za-z0-9_>./]*".
string name = 1;
// The operation name. There may be custom parameters in attrs.
// Op names starting with an underscore are reserved for internal use.
string op = 2;
// Each input is "node:src_output" with "node" being a string name and
// "src_output" indicating which output tensor to use from "node". If
// "src_output" is 0 the ":0" suffix can be omitted. Regular inputs
// may optionally be followed by control inputs that have the format
// "^node".
repeated string input = 3;
// A (possibly partial) specification for the device on which this
// node should be placed.
// The expected syntax for this string is as follows:
//
// DEVICE_SPEC ::= PARTIAL_SPEC
//
// PARTIAL_SPEC ::= ("/" CONSTRAINT) *
// CONSTRAINT ::= ("job:" JOB_NAME)
// | ("replica:" [1-9][0-9]*)
// | ("task:" [1-9][0-9]*)
// | ("device:" [A-Za-z]* ":" ([1-9][0-9]* | "*") )
//
// Valid values for this string include:
// * "/job:worker/replica:0/task:1/device:GPU:3" (full specification)
// * "/job:worker/device:GPU:3" (partial specification)
// * "" (no specification)
//
// If the constraints do not resolve to a single device (or if this
// field is empty or not present), the runtime will attempt to
// choose a device automatically.
string device = 4;
// Operation-specific graph-construction-time configuration.
// Note that this should include all attrs defined in the
// corresponding OpDef, including those with a value matching
// the default -- this allows the default to change and makes
// NodeDefs easier to interpret on their own. However, if
// an attr with a default is not specified in this list, the
// default will be used.
// The "names" (keys) must match the regexp "[a-z][a-z0-9_]+" (and
// one of the names from the corresponding OpDef's attr field).
// The values must have a type matching the corresponding OpDef
// attr's type field.
// TODO(josh11b): Add some examples here showing best practices.
map<string, AttrValue> attr = 5;
message ExperimentalDebugInfo {
// Opaque string inserted into error messages created by the runtime.
//
// This is intended to store the list of names of the nodes from the
// original graph that this node was derived. For example if this node, say
// C, was result of a fusion of 2 nodes A and B, then 'original_node' would
// be {A, B}. This information can be used to map errors originating at the
// current node to some top level source code.
repeated string original_node_names = 1;
// This is intended to store the list of names of the functions from the
// original graph that this node was derived. For example if this node, say
// C, was result of a fusion of node A in function FA and node B in function
// FB, then `original_funcs` would be {FA, FB}. If the node is in the top
// level graph, the `original_func` is empty. This information, with the
// `original_node_names` can be used to map errors originating at the
// current ndoe to some top level source code.
repeated string original_func_names = 2;
}
// This stores debug information associated with the node.
ExperimentalDebugInfo experimental_debug_info = 6;
}

View File

@ -0,0 +1,186 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "OpDefProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/op_def_go_proto";
import "attr_value.proto";
import "types.proto";
import "resource_handle.proto";
// Defines an operation. A NodeDef in a GraphDef specifies an Op by
// using the "op" field which should match the name of a OpDef.
// LINT.IfChange
message OpDef {
// Op names starting with an underscore are reserved for internal use.
// Names should be CamelCase and match the regexp "[A-Z][a-zA-Z0-9>_]*".
string name = 1;
// For describing inputs and outputs.
message ArgDef {
// Name for the input/output. Should match the regexp "[a-z][a-z0-9_]*".
string name = 1;
// Human readable description.
string description = 2;
// Describes the type of one or more tensors that are accepted/produced
// by this input/output arg. The only legal combinations are:
// * For a single tensor: either the "type" field is set or the
// "type_attr" field is set to the name of an attr with type "type".
// * For a sequence of tensors with the same type: the "number_attr"
// field will be set to the name of an attr with type "int", and
// either the "type" or "type_attr" field will be set as for
// single tensors.
// * For a sequence of tensors, the "type_list_attr" field will be set
// to the name of an attr with type "list(type)".
DataType type = 3;
string type_attr = 4; // if specified, attr must have type "type"
string number_attr = 5; // if specified, attr must have type "int"
// If specified, attr must have type "list(type)", and none of
// type, type_attr, and number_attr may be specified.
string type_list_attr = 6;
// The handle data for resource inputs.
repeated ResourceHandleProto.DtypeAndShape handle_data = 7;
// For inputs: if true, the inputs are required to be refs.
// By default, inputs can be either refs or non-refs.
// For outputs: if true, outputs are refs, otherwise they are not.
bool is_ref = 16;
};
// Description of the input(s).
repeated ArgDef input_arg = 2;
// Description of the output(s).
repeated ArgDef output_arg = 3;
// Named control outputs for this operation. Useful only for composite
// operations (i.e. functions) which want to name different control outputs.
repeated string control_output = 20;
// Description of the graph-construction-time configuration of this
// Op. That is to say, this describes the attr fields that will
// be specified in the NodeDef.
message AttrDef {
// A descriptive name for the argument. May be used, e.g. by the
// Python client, as a keyword argument name, and so should match
// the regexp "[a-z][a-z0-9_]+".
string name = 1;
// One of the type names from attr_value.proto ("string", "list(string)",
// "int", etc.).
string type = 2;
// A reasonable default for this attribute if the user does not supply
// a value. If not specified, the user must supply a value.
AttrValue default_value = 3;
// Human-readable description.
string description = 4;
// TODO(josh11b): bool is_optional?
// --- Constraints ---
// These constraints are only in effect if specified. Default is no
// constraints.
// For type == "int", this is a minimum value. For "list(___)"
// types, this is the minimum length.
bool has_minimum = 5;
int64 minimum = 6;
// The set of allowed values. Has type that is the "list" version
// of the "type" field above (uses the "list" field of AttrValue).
// If type == "type" or "list(type)" above, then the "type" field
// of "allowed_values.list" has the set of allowed DataTypes.
// If type == "string" or "list(string)", then the "s" field of
// "allowed_values.list" has the set of allowed strings.
AttrValue allowed_values = 7;
}
repeated AttrDef attr = 4;
// Optional deprecation based on GraphDef versions.
OpDeprecation deprecation = 8;
// One-line human-readable description of what the Op does.
string summary = 5;
// Additional, longer human-readable description of what the Op does.
string description = 6;
// -------------------------------------------------------------------------
// Which optimizations this operation can participate in.
// True if the operation is commutative ("op(a,b) == op(b,a)" for all inputs)
bool is_commutative = 18;
// If is_aggregate is true, then this operation accepts N >= 2
// inputs and produces 1 output all of the same type. Should be
// associative and commutative, and produce output with the same
// shape as the input. The optimizer may replace an aggregate op
// taking input from multiple devices with a tree of aggregate ops
// that aggregate locally within each device (and possibly within
// groups of nearby devices) before communicating.
// TODO(josh11b): Implement that optimization.
bool is_aggregate = 16; // for things like add
// Other optimizations go here, like
// can_alias_input, rewrite_when_output_unused, partitioning_strategy, etc.
// -------------------------------------------------------------------------
// Optimization constraints.
// Ops are marked as stateful if their behavior depends on some state beyond
// their input tensors (e.g. variable reading op) or if they have
// a side-effect (e.g. printing or asserting ops). Equivalently, stateless ops
// must always produce the same output for the same input and have
// no side-effects.
//
// By default Ops may be moved between devices. Stateful ops should
// either not be moved, or should only be moved if that state can also
// be moved (e.g. via some sort of save / restore).
// Stateful ops are guaranteed to never be optimized away by Common
// Subexpression Elimination (CSE).
bool is_stateful = 17; // for things like variables, queue
// -------------------------------------------------------------------------
// Non-standard options.
// By default, all inputs to an Op must be initialized Tensors. Ops
// that may initialize tensors for the first time should set this
// field to true, to allow the Op to take an uninitialized Tensor as
// input.
bool allows_uninitialized_input = 19; // for Assign, etc.
};
// LINT.ThenChange(
// https://www.tensorflow.org/code/tensorflow/core/framework/op_def_util.cc)
// Information about version-dependent deprecation of an op
message OpDeprecation {
// First GraphDef version at which the op is disallowed.
int32 version = 1;
// Explanation of why it was deprecated and what to use instead.
string explanation = 2;
};
// A collection of OpDefs
message OpList {
repeated OpDef op = 1;
};

View File

@ -0,0 +1,30 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "ReaderBaseProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/reader_base_go_proto";
// For serializing and restoring the state of ReaderBase, see
// reader_base.h for details.
message ReaderBaseState {
int64 work_started = 1;
int64 work_finished = 2;
int64 num_records_produced = 3;
bytes current_work = 4;
}

View File

@ -0,0 +1,60 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "graph.proto";
import "tensor_shape.proto";
import "types.proto";
option cc_enable_arenas = true;
option java_outer_classname = "RemoteFusedGraphExecuteInfoProto";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/remote_fused_graph_execute_info_go_proto";
// Protocol buffer representing a handle to a tensorflow resource. Handles are
// not valid across executions, but can be serialized back and forth from within
// a single run.
message RemoteFusedGraphExecuteInfo {
message TensorShapeTypeProto {
DataType dtype = 1;
TensorShapeProto shape = 2;
}
// Definition of remote graph
GraphDef remote_graph = 1;
// Remote fused graph input node name
repeated string graph_input_node_name = 2;
// Remote fused graph output node name
repeated string graph_output_node_name = 3;
// Executor's name
string executor_name = 4;
// Optional: Parameters given to the executor
bytes serialized_executor_parameters = 5;
// Optional: Default graph input tensor shape used to allocate memory
// before executing op
repeated TensorShapeTypeProto default_graph_input_tensor_shape = 6;
// Optional: Default graph input tensor shape used to allocate memory
// before executing op
// TODO(satok): Remote output tensor shape once shape information is stored
// in NodeDef
repeated TensorShapeTypeProto default_graph_output_tensor_shape = 7;
}

View File

@ -0,0 +1,57 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "tensor_shape.proto";
import "types.proto";
option cc_enable_arenas = true;
option java_outer_classname = "ResourceHandle";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/resource_handle_go_proto";
// Protocol buffer representing a handle to a tensorflow resource. Handles are
// not valid across executions, but can be serialized back and forth from within
// a single run.
message ResourceHandleProto {
// Unique name for the device containing the resource.
string device = 1;
// Container in which this resource is placed.
string container = 2;
// Unique name of this resource.
string name = 3;
// Hash code for the type of the resource. Is only valid in the same device
// and in the same execution.
uint64 hash_code = 4;
// For debug-only, the name of the type pointed to by this handle, if
// available.
string maybe_type_name = 5;
// Protocol buffer representing a pair of (data type, tensor shape).
message DtypeAndShape {
DataType dtype = 1;
TensorShapeProto shape = 2;
}
// Data types and shapes for the underlying resource.
repeated DtypeAndShape dtypes_and_shapes = 6;
reserved 7;
}

View File

@ -0,0 +1,100 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "allocation_description.proto";
import "tensor_description.proto";
option cc_enable_arenas = true;
option java_outer_classname = "StepStatsProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/step_stats_go_proto";
// An allocation/de-allocation operation performed by the allocator.
message AllocationRecord {
// The timestamp of the operation.
int64 alloc_micros = 1;
// Number of bytes allocated, or de-allocated if negative.
int64 alloc_bytes = 2;
}
message AllocatorMemoryUsed {
string allocator_name = 1;
// These are per-node allocator memory stats.
int64 total_bytes = 2;
int64 peak_bytes = 3;
// The bytes that are not deallocated.
int64 live_bytes = 4;
// The allocation and deallocation timeline.
repeated AllocationRecord allocation_records = 6;
// These are snapshots of the overall allocator memory stats.
// The number of live bytes currently allocated by the allocator.
int64 allocator_bytes_in_use = 5;
}
// Output sizes recorded for a single execution of a graph node.
message NodeOutput {
int32 slot = 1;
TensorDescription tensor_description = 3;
}
// For memory tracking.
message MemoryStats {
int64 temp_memory_size = 1;
int64 persistent_memory_size = 3;
repeated int64 persistent_tensor_alloc_ids = 5;
int64 device_temp_memory_size = 2 [deprecated = true];
int64 device_persistent_memory_size = 4 [deprecated = true];
repeated int64 device_persistent_tensor_alloc_ids = 6 [deprecated = true];
}
// Time/size stats recorded for a single execution of a graph node.
message NodeExecStats {
// TODO(tucker): Use some more compact form of node identity than
// the full string name. Either all processes should agree on a
// global id (cost_id?) for each node, or we should use a hash of
// the name.
string node_name = 1;
int64 all_start_micros = 2;
int64 op_start_rel_micros = 3;
int64 op_end_rel_micros = 4;
int64 all_end_rel_micros = 5;
repeated AllocatorMemoryUsed memory = 6;
repeated NodeOutput output = 7;
string timeline_label = 8;
int64 scheduled_micros = 9;
uint32 thread_id = 10;
repeated AllocationDescription referenced_tensor = 11;
MemoryStats memory_stats = 12;
int64 all_start_nanos = 13;
int64 op_start_rel_nanos = 14;
int64 op_end_rel_nanos = 15;
int64 all_end_rel_nanos = 16;
int64 scheduled_nanos = 17;
}
message DeviceStepStats {
string device = 1;
repeated NodeExecStats node_stats = 2;
// Its key is thread id.
map<uint32, string> thread_names = 3;
}
message StepStats {
repeated DeviceStepStats dev_stats = 1;
}

View File

@ -0,0 +1,161 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "tensor.proto";
option cc_enable_arenas = true;
option java_outer_classname = "SummaryProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/summary_go_proto";
// Metadata associated with a series of Summary data
message SummaryDescription {
// Hint on how plugins should process the data in this series.
// Supported values include "scalar", "histogram", "image", "audio"
string type_hint = 1;
}
// Serialization format for histogram module in
// core/lib/histogram/histogram.h
message HistogramProto {
double min = 1;
double max = 2;
double num = 3;
double sum = 4;
double sum_squares = 5;
// Parallel arrays encoding the bucket boundaries and the bucket values.
// bucket(i) is the count for the bucket i. The range for
// a bucket is:
// i == 0: -DBL_MAX .. bucket_limit(0)
// i != 0: bucket_limit(i-1) .. bucket_limit(i)
repeated double bucket_limit = 6 [packed = true];
repeated double bucket = 7 [packed = true];
}
// A SummaryMetadata encapsulates information on which plugins are able to make
// use of a certain summary value.
message SummaryMetadata {
message PluginData {
// The name of the plugin this data pertains to.
string plugin_name = 1;
// The content to store for the plugin. The best practice is for this to be
// a binary serialized protocol buffer.
bytes content = 2;
}
// Data that associates a summary with a certain plugin.
PluginData plugin_data = 1;
// Display name for viewing in TensorBoard.
string display_name = 2;
// Longform readable description of the summary sequence. Markdown supported.
string summary_description = 3;
// Class of data stored in this time series. Required for compatibility with
// TensorBoard's generic data facilities (`DataProvider`, et al.). This value
// imposes constraints on the dtype and shape of the corresponding tensor
// values. See `DataClass` docs for details.
DataClass data_class = 4;
}
enum DataClass {
// Unknown data class, used (implicitly) for legacy data. Will not be
// processed by data ingestion pipelines.
DATA_CLASS_UNKNOWN = 0;
// Scalar time series. Each `Value` for the corresponding tag must have
// `tensor` set to a rank-0 tensor of type `DT_FLOAT` (float32).
DATA_CLASS_SCALAR = 1;
// Tensor time series. Each `Value` for the corresponding tag must have
// `tensor` set. The tensor value is arbitrary, but should be small to
// accommodate direct storage in database backends: an upper bound of a few
// kilobytes is a reasonable rule of thumb.
DATA_CLASS_TENSOR = 2;
// Blob sequence time series. Each `Value` for the corresponding tag must
// have `tensor` set to a rank-1 tensor of bytestring dtype.
DATA_CLASS_BLOB_SEQUENCE = 3;
}
// A Summary is a set of named values to be displayed by the
// visualizer.
//
// Summaries are produced regularly during training, as controlled by
// the "summary_interval_secs" attribute of the training operation.
// Summaries are also produced at the end of an evaluation.
message Summary {
message Image {
// Dimensions of the image.
int32 height = 1;
int32 width = 2;
// Valid colorspace values are
// 1 - grayscale
// 2 - grayscale + alpha
// 3 - RGB
// 4 - RGBA
// 5 - DIGITAL_YUV
// 6 - BGRA
int32 colorspace = 3;
// Image data in encoded format. All image formats supported by
// image_codec::CoderUtil can be stored here.
bytes encoded_image_string = 4;
}
message Audio {
// Sample rate of the audio in Hz.
float sample_rate = 1;
// Number of channels of audio.
int64 num_channels = 2;
// Length of the audio in frames (samples per channel).
int64 length_frames = 3;
// Encoded audio data and its associated RFC 2045 content type (e.g.
// "audio/wav").
bytes encoded_audio_string = 4;
string content_type = 5;
}
message Value {
// This field is deprecated and will not be set.
string node_name = 7;
// Tag name for the data. Used by TensorBoard plugins to organize data. Tags
// are often organized by scope (which contains slashes to convey
// hierarchy). For example: foo/bar/0
string tag = 1;
// Contains metadata on the summary value such as which plugins may use it.
// Take note that many summary values may lack a metadata field. This is
// because the FileWriter only keeps a metadata object on the first summary
// value with a certain tag for each tag. TensorBoard then remembers which
// tags are associated with which plugins. This saves space.
SummaryMetadata metadata = 9;
// Value associated with the tag.
oneof value {
float simple_value = 2;
bytes obsolete_old_style_histogram = 3;
Image image = 4;
HistogramProto histo = 5;
Audio audio = 6;
TensorProto tensor = 8;
}
}
// Set of values for the summary.
repeated Value value = 1;
}

View File

@ -0,0 +1,108 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "resource_handle.proto";
import "tensor_shape.proto";
import "types.proto";
option cc_enable_arenas = true;
option java_outer_classname = "TensorProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/tensor_go_proto";
// Protocol buffer representing a tensor.
message TensorProto {
DataType dtype = 1;
// Shape of the tensor. TODO(touts): sort out the 0-rank issues.
TensorShapeProto tensor_shape = 2;
// Only one of the representations below is set, one of "tensor_contents" and
// the "xxx_val" attributes. We are not using oneof because as oneofs cannot
// contain repeated fields it would require another extra set of messages.
// Version number.
//
// In version 0, if the "repeated xxx" representations contain only one
// element, that element is repeated to fill the shape. This makes it easy
// to represent a constant Tensor with a single value.
int32 version_number = 3;
// Serialized raw tensor content from either Tensor::AsProtoTensorContent or
// memcpy in tensorflow::grpc::EncodeTensorToByteBuffer. This representation
// can be used for all tensor types. The purpose of this representation is to
// reduce serialization overhead during RPC call by avoiding serialization of
// many repeated small items.
bytes tensor_content = 4;
// Type specific representations that make it easy to create tensor protos in
// all languages. Only the representation corresponding to "dtype" can
// be set. The values hold the flattened representation of the tensor in
// row major order.
// DT_HALF, DT_BFLOAT16. Note that since protobuf has no int16 type, we'll
// have some pointless zero padding for each value here.
repeated int32 half_val = 13 [packed = true];
// DT_FLOAT.
repeated float float_val = 5 [packed = true];
// DT_DOUBLE.
repeated double double_val = 6 [packed = true];
// DT_INT32, DT_INT16, DT_INT8, DT_UINT8.
repeated int32 int_val = 7 [packed = true];
// DT_STRING
repeated bytes string_val = 8;
// DT_COMPLEX64. scomplex_val(2*i) and scomplex_val(2*i+1) are real
// and imaginary parts of i-th single precision complex.
repeated float scomplex_val = 9 [packed = true];
// DT_INT64
repeated int64 int64_val = 10 [packed = true];
// DT_BOOL
repeated bool bool_val = 11 [packed = true];
// DT_COMPLEX128. dcomplex_val(2*i) and dcomplex_val(2*i+1) are real
// and imaginary parts of i-th double precision complex.
repeated double dcomplex_val = 12 [packed = true];
// DT_RESOURCE
repeated ResourceHandleProto resource_handle_val = 14;
// DT_VARIANT
repeated VariantTensorDataProto variant_val = 15;
// DT_UINT32
repeated uint32 uint32_val = 16 [packed = true];
// DT_UINT64
repeated uint64 uint64_val = 17 [packed = true];
}
// Protocol buffer representing the serialization format of DT_VARIANT tensors.
message VariantTensorDataProto {
// Name of the type of objects being serialized.
string type_name = 1;
// Portions of the object that are not Tensors.
bytes metadata = 2;
// Tensors contained within objects being serialized.
repeated TensorProto tensors = 3;
}

View File

@ -0,0 +1,36 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
import "allocation_description.proto";
import "tensor_shape.proto";
import "types.proto";
option cc_enable_arenas = true;
option java_outer_classname = "TensorDescriptionProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/tensor_description_go_proto";
message TensorDescription {
// Data type of tensor elements
DataType dtype = 1;
// Shape of the tensor.
TensorShapeProto shape = 2;
// Information about the size and allocator used for the data
AllocationDescription allocation_description = 4;
}

View File

@ -0,0 +1,58 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
// Protocol buffer representing the shape of tensors.
syntax = "proto3";
option cc_enable_arenas = true;
option java_outer_classname = "TensorShapeProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/tensor_shape_go_proto";
package tensorflow;
// Dimensions of a tensor.
message TensorShapeProto {
// One dimension of the tensor.
message Dim {
// Size of the tensor in that dimension.
// This value must be >= -1, but values of -1 are reserved for "unknown"
// shapes (values of -1 mean "unknown" dimension). Certain wrappers
// that work with TensorShapeProto may fail at runtime when deserializing
// a TensorShapeProto containing a dim value of -1.
int64 size = 1;
// Optional name of the tensor dimension.
string name = 2;
};
// Dimensions of the tensor, such as {"input", 30}, {"output", 40}
// for a 30 x 40 2D tensor. If an entry has size -1, this
// corresponds to a dimension of unknown size. The names are
// optional.
//
// The order of entries in "dim" matters: It indicates the layout of the
// values in the tensor in-memory representation.
//
// The first entry in "dim" is the outermost dimension used to layout the
// values, the last entry is the innermost dimension. This matches the
// in-memory layout of RowMajor Eigen tensors.
//
// If "dim.size()" > 0, "unknown_rank" must be false.
repeated Dim dim = 2;
// If true, the number of dimensions in the shape is unknown.
//
// If true, "dim.size()" must be 0.
bool unknown_rank = 3;
};

View File

@ -0,0 +1,51 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
// Protocol buffer representing slices of a tensor
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "TensorSliceProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/tensor_slice_go_proto";
// Can only be interpreted if you know the corresponding TensorShape.
message TensorSliceProto {
// Extent of the slice in one dimension.
message Extent {
// Either both or no attributes must be set. When no attribute is set
// means: All data in that dimension.
// Start index of the slice, starting at 0.
int64 start = 1;
// Length of the slice: if the length is missing or -1 we will
// interpret this as "everything in this dimension". We use
// "oneof" to preserve information about whether the length is
// present without changing the serialization format from the
// prior proto2 version of this proto.
oneof has_length {
int64 length = 2;
}
}
// Extent of the slice in all tensor dimensions.
//
// Must have one entry for each of the dimension of the tensor that this
// slice belongs to. The order of sizes is the same as the order of
// dimensions in the TensorShape.
repeated Extent extent = 1;
}

View File

@ -0,0 +1,101 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "TypesProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/types_go_proto";
// (== suppress_warning documentation-presence ==)
// LINT.IfChange
enum DataType {
// Not a legal value for DataType. Used to indicate a DataType field
// has not been set.
DT_INVALID = 0;
// Data types that all computation devices are expected to be
// capable to support.
DT_FLOAT = 1;
DT_DOUBLE = 2;
DT_INT32 = 3;
DT_UINT8 = 4;
DT_INT16 = 5;
DT_INT8 = 6;
DT_STRING = 7;
DT_COMPLEX64 = 8; // Single-precision complex
DT_INT64 = 9;
DT_BOOL = 10;
DT_QINT8 = 11; // Quantized int8
DT_QUINT8 = 12; // Quantized uint8
DT_QINT32 = 13; // Quantized int32
DT_BFLOAT16 = 14; // Float32 truncated to 16 bits. Only for cast ops.
DT_QINT16 = 15; // Quantized int16
DT_QUINT16 = 16; // Quantized uint16
DT_UINT16 = 17;
DT_COMPLEX128 = 18; // Double-precision complex
DT_HALF = 19;
DT_RESOURCE = 20;
DT_VARIANT = 21; // Arbitrary C++ data types
DT_UINT32 = 22;
DT_UINT64 = 23;
// Do not use! These are only for parameters. Every enum above
// should have a corresponding value below (verified by types_test).
DT_FLOAT_REF = 101;
DT_DOUBLE_REF = 102;
DT_INT32_REF = 103;
DT_UINT8_REF = 104;
DT_INT16_REF = 105;
DT_INT8_REF = 106;
DT_STRING_REF = 107;
DT_COMPLEX64_REF = 108;
DT_INT64_REF = 109;
DT_BOOL_REF = 110;
DT_QINT8_REF = 111;
DT_QUINT8_REF = 112;
DT_QINT32_REF = 113;
DT_BFLOAT16_REF = 114;
DT_QINT16_REF = 115;
DT_QUINT16_REF = 116;
DT_UINT16_REF = 117;
DT_COMPLEX128_REF = 118;
DT_HALF_REF = 119;
DT_RESOURCE_REF = 120;
DT_VARIANT_REF = 121;
DT_UINT32_REF = 122;
DT_UINT64_REF = 123;
}
// LINT.ThenChange(
// https://www.tensorflow.org/code/tensorflow/c/tf_datatype.h,
// https://www.tensorflow.org/code/tensorflow/go/tensor.go,
// https://www.tensorflow.org/code/tensorflow/core/framework/tensor.cc,
// https://www.tensorflow.org/code/tensorflow/core/framework/types.h,
// https://www.tensorflow.org/code/tensorflow/core/framework/types.cc,
// https://www.tensorflow.org/code/tensorflow/python/framework/dtypes.py,
// https://www.tensorflow.org/code/tensorflow/python/framework/function.py)
// For identifying the underlying type of a variant. For variants, the types
// listed here are a subset of the types in the variant type registry,
// corresponding to commonly used variants which must occasionally be
// special-cased.
enum SpecializedType {
// Invalid/unknown specialized type.
ST_INVALID = 0;
// "tensorflow::TensorList" in the variant type registry.
ST_TENSOR_LIST = 1;
// "tensorflow::data::Optional" in the variant type registry.
ST_OPTIONAL = 2;
}

View File

@ -0,0 +1,96 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "VariableProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/variable_go_proto";
// Indicates when a distributed variable will be synced.
enum VariableSynchronization {
// `AUTO`: Indicates that the synchronization will be determined by the
// current `DistributionStrategy` (eg. With `MirroredStrategy` this would be
// `ON_WRITE`).
VARIABLE_SYNCHRONIZATION_AUTO = 0;
// `NONE`: Indicates that there will only be one copy of the variable, so
// there is no need to sync.
VARIABLE_SYNCHRONIZATION_NONE = 1;
// `ON_WRITE`: Indicates that the variable will be updated across devices
// every time it is written.
VARIABLE_SYNCHRONIZATION_ON_WRITE = 2;
// `ON_READ`: Indicates that the variable will be aggregated across devices
// when it is read (eg. when checkpointing or when evaluating an op that uses
// the variable).
VARIABLE_SYNCHRONIZATION_ON_READ = 3;
}
// Indicates how a distributed variable will be aggregated.
enum VariableAggregation {
// `NONE`: This is the default, giving an error if you use a
// variable-update operation with multiple replicas.
VARIABLE_AGGREGATION_NONE = 0;
// `SUM`: Add the updates across replicas.
VARIABLE_AGGREGATION_SUM = 1;
// `MEAN`: Take the arithmetic mean ("average") of the updates across
// replicas.
VARIABLE_AGGREGATION_MEAN = 2;
// `ONLY_FIRST_REPLICA`: This is for when every replica is performing the same
// update, but we only want to perform the update once. Used, e.g., for the
// global step counter.
VARIABLE_AGGREGATION_ONLY_FIRST_REPLICA = 3;
}
// Protocol buffer representing a Variable.
message VariableDef {
// Name of the variable tensor.
string variable_name = 1;
// Name of the tensor holding the variable's initial value.
string initial_value_name = 6;
// Name of the initializer op.
string initializer_name = 2;
// Name of the snapshot tensor.
string snapshot_name = 3;
// Support for saving variables as slices of a larger variable.
SaveSliceInfoDef save_slice_info_def = 4;
// Whether to represent this as a ResourceVariable.
bool is_resource = 5;
// Whether this variable should be trained.
bool trainable = 7;
// Indicates when a distributed variable will be synced.
VariableSynchronization synchronization = 8;
// Indicates how a distributed variable will be aggregated.
VariableAggregation aggregation = 9;
}
message SaveSliceInfoDef {
// Name of the full variable of which this is a slice.
string full_name = 1;
// Shape of the full variable.
repeated int64 full_shape = 2;
// Offset of this variable into the full variable.
repeated int64 var_offset = 3;
// Shape of this variable.
repeated int64 var_shape = 4;
}

View File

@ -0,0 +1,45 @@
/* Copyright 2015 The TensorFlow Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.*/
// Modification Copyright (C) 2021 Intel Corporation
syntax = "proto3";
package tensorflow;
option cc_enable_arenas = true;
option java_outer_classname = "VersionsProtos";
option java_multiple_files = true;
option java_package = "org.tensorflow.framework";
option go_package = "github.com/tensorflow/tensorflow/tensorflow/go/core/framework/versions_go_proto";
// Version information for a piece of serialized data
//
// There are different types of versions for each type of data
// (GraphDef, etc.), but they all have the same common shape
// described here.
//
// Each consumer has "consumer" and "min_producer" versions (specified
// elsewhere). A consumer is allowed to consume this data if
//
// producer >= min_producer
// consumer >= min_consumer
// consumer not in bad_consumers
//
message VersionDef {
// The version of the code that produced this data.
int32 producer = 1;
// Any consumer below this version is not allowed to consume this data.
int32 min_consumer = 2;
// Specific consumer versions which are disallowed (e.g. due to bugs).
repeated int32 bad_consumers = 3;
}

View File

@ -0,0 +1,19 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <frontend_manager/frontend_manager.hpp>
#include <tensorflow_frontend/frontend.hpp>
extern "C" NGRAPH_HELPER_DLL_EXPORT ngraph::frontend::FrontEndVersion GetAPIVersion() {
return OV_FRONTEND_API_VERSION;
}
extern "C" NGRAPH_HELPER_DLL_EXPORT void* GetFrontEndData() {
auto res = new ngraph::frontend::FrontEndPluginInfo();
res->m_name = "tf";
res->m_creator = []() {
return std::make_shared<ov::frontend::FrontEndTF>();
};
return res;
}

View File

@ -0,0 +1,18 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <tf_framework_node.hpp>
namespace ov {
namespace frontend {
namespace tf {
void TFFrameworkNode::validate_and_infer_types() {
for (size_t i = 0; i < get_output_size(); ++i) {
set_output_type(i, ov::element::dynamic, PartialShape::dynamic());
}
}
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,50 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#pragma once
#include <algorithm>
#include <openvino/op/util/framework_node.hpp>
#include <tensorflow_frontend/place.hpp>
#include "graph_iterator_proto.hpp"
namespace ov {
namespace frontend {
namespace tf {
class TFFrameworkNode : public ::ov::op::util::FrameworkNode {
public:
OPENVINO_OP("TFFrameworkNode", "util", ::ov::op::util::FrameworkNode);
TFFrameworkNode(const std::shared_ptr<DecoderBase>& decoder, const OutputVector& inputs, size_t num_outputs)
: FrameworkNode(inputs, std::max(num_outputs, size_t(1))),
m_decoder(decoder) {
ov::op::util::FrameworkNodeAttrs attrs;
attrs.set_type_name(m_decoder->get_op_type());
set_attrs(attrs);
validate_and_infer_types();
}
void validate_and_infer_types() override;
std::shared_ptr<Node> clone_with_new_inputs(const OutputVector& inputs) const override {
return std::make_shared<TFFrameworkNode>(m_decoder, inputs, get_output_size());
}
std::string get_op_type() const {
return m_decoder->get_op_type();
}
std::shared_ptr<DecoderBase> get_decoder() const {
return m_decoder;
}
private:
std::shared_ptr<DecoderBase> m_decoder;
};
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -0,0 +1,20 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include "utils.hpp"
void ov::frontend::tf::SetTracingInfo(const std::string& op_name, const ov::Output<ov::Node>& ng_node) {
auto node = ng_node.get_node_shared_ptr();
node->set_friendly_name(op_name);
node->add_provenance_tag(op_name);
}
void ov::frontend::tf::TFTensorShapeToNGraphShape(const tensorflow::TensorShapeProto& tf_shape,
ov::PartialShape* ng_shape) {
std::vector<ov::Dimension> dims;
for (int i = 0; i < tf_shape.dim_size(); i++) {
dims.push_back(tf_shape.dim(i).size());
}
*ng_shape = ov::PartialShape(dims);
}

View File

@ -0,0 +1,381 @@
/* Copyright (C) 2018-2021 Intel Corporation
* SPDX-License-Identifier: Apache-2.0
*
* Copyright 2017 The TensorFlow Authors. All Rights Reserved.
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
* http://www.apache.org/licenses/LICENSE-2.0
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*
* We modified "ValuesFromConstNode" function from
* tensorflow/core/grappler/optimizers/arithmetic_optimizer.cc file
* to integrate it with our infrastructure. The purpose and basic
* functionality remains the same.
==============================================================================*/
#pragma once
#include <openvino/opsets/opset8.hpp>
#include "graph_iterator_proto.hpp"
#include "ngraph/ngraph.hpp"
#include "ngraph_conversions.hpp"
#include "node_context.hpp"
namespace ov {
namespace frontend {
namespace tf {
namespace detail {
// TODO: avoid using directly:
using ::tensorflow::DataType;
using ::tensorflow::TensorProto;
// TODO: separate interface from proto implementation; here is a proto implementation
class TensorWrapper {
public:
const TensorProto* tensor_def;
TensorWrapper(const TensorProto* _tensor_def) : tensor_def(_tensor_def) {}
// a hack to minimize amount of code
TensorWrapper& attrs() const {
return const_cast<TensorWrapper&>(*this);
}
template <typename T>
std::vector<T> flat() const;
size_t NumElements() const;
DataType dtype() const;
};
} // namespace detail
} // namespace tf
} // namespace frontend
} // namespace ov
namespace ov {
namespace frontend {
namespace tf {
using namespace ::ov::frontend::tf::detail;
using OpMap = std::unordered_map<std::string, std::vector<ov::Output<ov::Node>>>;
void extract_operation_name_and_port(const std::string& port_name,
std::string& operation_name,
size_t& port_index,
std::string& port_type);
class Status {
public:
int status = 0;
std::string message;
static Status OK() {
return Status();
}
Status(const std::string& x) : message(x), status(1) {}
Status() {}
};
inline bool operator!=(const Status& x, const Status& y) {
return x.status != y.status;
}
inline std::ostream& operator<<(std::ostream& out, const Status& s) {
return out << s.message;
}
#define TF_RETURN_IF_ERROR(S) \
if ((S).status != 0) \
throw S;
class errors {
public:
static Status InvalidArgument(const std::string& x) {
return Status("InvalidArgument: " + x);
}
static Status Internal(const std::string& x) {
return Status("Internal: " + x);
}
static Status Unimplemented(const std::string& x) {
return Status("Unimplemented: " + x);
}
};
void SetTracingInfo(const std::string& op_name, const ov::Output<ov::Node>& ng_node);
template <typename T>
static void MakePadding(const std::string& tf_padding_type,
const ov::Shape& ng_image_shape,
const ov::Shape& ng_kernel_shape,
const ov::Strides& ng_strides,
const ov::Shape& ng_dilations,
T& ng_padding_below,
T& ng_padding_above) {
if (tf_padding_type == "SAME") {
ov::Shape img_shape = {0, 0};
img_shape.insert(img_shape.end(), ng_image_shape.begin(), ng_image_shape.end());
ov::infer_auto_padding(img_shape,
ng_kernel_shape,
ng_strides,
ng_dilations,
ov::op::PadType::SAME_UPPER,
ng_padding_above,
ng_padding_below);
} else if (tf_padding_type == "VALID") {
ng_padding_below.assign(ng_image_shape.size(), 0);
ng_padding_above.assign(ng_image_shape.size(), 0);
}
}
template <typename Ttensor, typename Tvector>
static void ConvertTensorDataToVector(const TensorWrapper& tensor, std::vector<Tvector>* vector) {
const Ttensor* data = tensor.flat<Ttensor>().data();
vector->resize(tensor.NumElements());
for (int64_t i = 0; i < tensor.NumElements(); i++) {
(*vector)[i] = Tvector(data[i]);
}
}
static bool VecStrCmp(const std::vector<std::string>& a, const std::vector<std::string>& b) {
return a == b;
}
static Status ValidateInputCount(const NodeContext& op, size_t count) {
if (op.get_ng_input_size() != count) {
std::ostringstream buf;
buf << "\"" << op.get_name() << "\" requires " << count << " input(s), got " << op.get_ng_input_size()
<< " instead";
return errors::InvalidArgument(buf.str());
}
return Status::OK();
}
static void ValidateInputCountMin(const ov::frontend::tf::NodeContext& node, size_t count) {
if (node.get_ng_input_size() < count) {
std::ostringstream buf;
buf << "\"" << node.get_name() << "\" requires at least " << count << " input(s), got "
<< node.get_ng_input_size() << " instead";
throw errors::InvalidArgument(buf.str());
}
}
// Check to make sure the axis dimension for reduction are in within range.
// Returns error if axis is out of range. Otherwise returns Status::OK().
static Status CheckAxisDimInRange(const std::vector<int64_t>& axes, size_t rank) {
for (auto i : axes) {
if (i < -(int)rank || i >= (int)rank) {
std::ostringstream buf;
buf << "Axis Dimension is out of range. Got " << i << ", should be in range [-" << rank << ", " << rank
<< ")";
return errors::InvalidArgument(buf.str());
}
}
return Status::OK();
}
//
// Helper for storing ops in ng_op_map.
// For most of the cases, op would have one output so
// std::vector ng_op_map[op_name] would contain one element.
//
// If storing more than one output_nodes, make sure it's in
// the same order as tensorflow would do that.
//
// Parameters:
// Builder::OpMap& ng_op_map - The TF-to-nGraph op map.
// std::string op_name - Name of the op.
//
// ov::Output<ov::Node> output_node - ov::Node to store
//
static void SaveNgOp(OpMap& ng_op_map, const std::string& op_name, ov::Output<ov::Node> output_node) {
// no need to try-catch, map[key] will create std::vector object
// if not exists
ng_op_map[op_name].push_back(output_node);
}
template <class TOpType, class... TArg>
ov::Output<ov::Node> ConstructNgNode(const std::string& op_name, TArg&&... Args) {
auto ng_node = std::make_shared<TOpType>(std::forward<TArg>(Args)...);
SetTracingInfo(op_name, ng_node);
return ng_node;
}
static Status GetInputNode(const NodeContext op, size_t input_idx, ov::Output<ov::Node>& result) {
// Stub
result = op.get_ng_input(input_idx);
return Status::OK();
}
namespace detail {
static Status GetInputNodes(const NodeContext&, size_t) {
return Status::OK();
}
template <typename... Arguments>
static Status GetInputNodes(const NodeContext node,
size_t index,
ov::Output<ov::Node>& result,
Arguments&... remaining) {
TF_RETURN_IF_ERROR(GetInputNode(node, index, result));
return GetInputNodes(node, index + 1, remaining...);
}
} // namespace detail
template <typename... Arguments>
static Status GetInputNodes(const NodeContext& node, Arguments&... remaining) {
constexpr size_t args_len = sizeof...(Arguments);
TF_RETURN_IF_ERROR(ValidateInputCount(node, args_len));
return detail::GetInputNodes(node, 0, remaining...);
}
void TFTensorShapeToNGraphShape(const ::tensorflow::TensorShapeProto& tf_shape, ov::PartialShape* ng_shape);
template <typename T>
static void GetStaticInputVector(const ov::frontend::tf::NodeContext& node,
int64_t input_index,
std::vector<T>* vector) {
ov::Output<ov::Node> ng_input = node.get_ng_input(input_index);
auto constant = std::dynamic_pointer_cast<ov::opset8::Constant>(ng_input.get_node_shared_ptr());
FRONT_END_GENERAL_CHECK(constant != nullptr, "Node ", node.get_name(), " can't be casted to Constant.");
*vector = constant->cast_vector<T>();
}
// Taken from: tensorflow/core/grappler/optimizers/arithmetic_optimizer.cc
// Extract values from a Const op to `values`. Returns true if succeeds.
//
// Modified with an extra `VecT` parameter to handle the case where the type
// in the std::vector does not match TensorFlow's notion of what the C++ type
// should be (e.g. when T is `bool`, we actually need a std::vector of `char` for
// compatibility with nGraph).
template <typename T, typename VecT = T>
static Status ValuesFromConstNode(const ::ov::frontend::DecoderBase* node,
ov::Shape* const_tensor_shape,
std::vector<VecT>* values) {
if (node->get_op_type() != "Const") {
return errors::InvalidArgument("TFNodeDecoder not a Const");
}
auto dt1 = node->get_attribute("dtype", ::ov::VariantWrapper<::tensorflow::DataType>::get_type_info_static());
FRONT_END_GENERAL_CHECK(dt1);
auto dt = std::dynamic_pointer_cast<::ov::VariantWrapper<::tensorflow::DataType>>(dt1)->get();
auto tensor_proto_var =
node->get_attribute("value", ::ov::VariantWrapper<::tensorflow::TensorProto>::get_type_info_static());
FRONT_END_GENERAL_CHECK(tensor_proto_var);
auto tensor_proto =
std::dynamic_pointer_cast<::ov::VariantWrapper<::tensorflow::TensorProto>>(tensor_proto_var)->get();
const tensorflow::TensorShapeProto& shape = tensor_proto.tensor_shape();
ov::PartialShape pshape;
TFTensorShapeToNGraphShape(shape, &pshape);
*const_tensor_shape = pshape.get_shape();
FRONT_END_GENERAL_CHECK(!pshape.is_dynamic(), "Dynamic shapes are not supported in ValuesFromConstNode function");
auto tensor_content = tensor_proto.tensor_content();
std::vector<char> tensor_values_plain(tensor_content.begin(), tensor_content.end());
const T* tensor_values = reinterpret_cast<const T*>(tensor_values_plain.data());
if (!tensor_values_plain.empty() && tensor_proto.has_tensor_shape()) {
// When tensor_shape is set, theoretically the representation of the data
// could be compressed. So, before copying values to the returned vector,
// make sure no compression happens.
// if (shape.dim_size() == 1 && shape.dim(0).size() == tensor_values_plain.size()/sizeof(T)) {
values->insert(values->end(), tensor_values, tensor_values + tensor_values_plain.size() / sizeof(T));
return Status::OK();
}
const auto tensor_content_size = tensor_proto.tensor_content().size();
if (tensor_content_size % sizeof(VecT)) {
std::cerr << "[ ERROR ] tensor_content_size (" << tensor_content_size << ") is not a multiple of "
<< sizeof(VecT);
}
// If tensor_content_size is zero, we'll have to take the values from
// int_val, float_val, etc.
if (tensor_content_size == 0) {
int64_t n_elements = 1;
for (auto i = 0; i < shape.dim_size(); i++) {
if (shape.dim(i).size() < 0) {
return errors::InvalidArgument("Const node has empty tensor and an unknown dimension size");
}
n_elements *= shape.dim(i).size();
}
values->resize(n_elements);
auto val_lastsaved = (T)0; // cast
for (auto i = 0; i < n_elements; i++) {
int64_t val_size = 0;
auto val_i = (T)0; // cast
switch (dt) {
// TODO: there are more element types to support
// here
case tensorflow::DT_INT32:
val_size = tensor_proto.int_val_size();
if (val_size > 0)
val_i = tensor_proto.int_val()[i];
break;
case tensorflow::DT_INT64:
val_size = tensor_proto.int64_val_size();
if (val_size > 0)
val_i = tensor_proto.int64_val()[i];
break;
case tensorflow::DT_FLOAT:
val_size = tensor_proto.float_val_size();
if (val_size > 0)
val_i = tensor_proto.float_val()[i];
break;
case tensorflow::DT_BOOL:
val_size = tensor_proto.bool_val_size();
if (val_size > 0)
val_i = tensor_proto.bool_val()[i];
break;
case tensorflow::DT_DOUBLE:
val_size = tensor_proto.double_val_size();
if (val_size > 0)
val_i = tensor_proto.double_val()[i];
break;
default:
NGRAPH_VLOG(0) << "Const node has empty tensor_proto and we don't know how to "
"handle this element type";
return errors::Unimplemented("Encountered unknown element type " + DataType_Name(dt) +
" on an empty tensor_proto");
}
if (val_size == 0) {
return errors::InvalidArgument("Empty values vector");
} else if (i < val_size) {
(*values)[i] = val_i;
val_lastsaved = val_i;
} else {
(*values)[i] = val_lastsaved;
}
}
} else {
return Status::OK();
}
return Status::OK();
}
template <typename T, typename VecT = T>
static Status MakeConstOp(const NodeContext& node, ov::element::Type et, ov::Output<ov::Node>& ng_node) {
std::vector<VecT> const_values;
ov::Shape ng_shape;
TF_RETURN_IF_ERROR((ValuesFromConstNode<T, VecT>(node.get_decoder(), &ng_shape, &const_values)));
ng_node = ConstructNgNode<ov::opset8::Constant>(node.get_name(), et, ng_shape, const_values);
return Status::OK();
}
} // namespace tf
} // namespace frontend
} // namespace ov

View File

@ -11,6 +11,10 @@ if (NGRAPH_ONNX_FRONTEND_ENABLE)
add_subdirectory(onnx)
endif()
if (NGRAPH_TF_FRONTEND_ENABLE)
add_subdirectory(tensorflow)
endif()
set(SRC ${CMAKE_CURRENT_SOURCE_DIR}/mock_frontend.cpp)
add_library(mock1_ngraph_frontend SHARED ${SRC})

View File

@ -0,0 +1,66 @@
# Copyright (C) 2018-2021 Intel Corporation
# SPDX-License-Identifier: Apache-2.0
#
set(TARGET_NAME "tensorflow_tests")
file(GLOB_RECURSE SRC ${CMAKE_CURRENT_SOURCE_DIR}/*.cpp)
add_executable(${TARGET_NAME} ${SRC})
target_link_libraries(${TARGET_NAME} PRIVATE frontend_shared_test_classes)
add_clang_format_target(${TARGET_NAME}_clang FOR_TARGETS ${TARGET_NAME})
install(TARGETS ${TARGET_NAME}
RUNTIME DESTINATION tests
COMPONENT tests
EXCLUDE_FROM_ALL)
# Test model generating
ie_check_pip_package(tensorflow WARNING)
set(TEST_TENSORFLOW_MODELS_DIRNAME test_model_zoo/tensorflow_test_models)
target_compile_definitions(${TARGET_NAME} PRIVATE -D TEST_TENSORFLOW_MODELS_DIRNAME=\"${TEST_TENSORFLOW_MODELS_DIRNAME}/\")
# If 'tensorflow' is not found, code will still be compiled
# but models will not be generated and tests will fail
# This is done this way for 'code style' and check cases - cmake shall pass, but CI machine doesn't need to have
# 'tensorflow' installed to check code style
if (tensorflow_FOUND)
set(TEST_TENSORFLOW_MODELS ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/${TEST_TENSORFLOW_MODELS_DIRNAME}/)
file(GLOB_RECURSE TENSORFLOW_GEN_SCRIPTS ${CMAKE_CURRENT_SOURCE_DIR}/test_models/gen_scripts/generate_*.py)
file(GLOB_RECURSE TENSORFLOW_ALL_SCRIPTS ${CMAKE_CURRENT_SOURCE_DIR}/*.py)
set(OUT_FILES "")
foreach(GEN_SCRIPT ${TENSORFLOW_GEN_SCRIPTS})
get_filename_component(FILE_WE ${GEN_SCRIPT} NAME_WE)
set(OUT_DONE_FILE ${TEST_TENSORFLOW_MODELS}/${FILE_WE}_done.txt)
set(OUT_FILES ${OUT_DONE_FILE} ${OUT_FILES})
add_custom_command(OUTPUT ${OUT_DONE_FILE}
COMMAND ${PYTHON_EXECUTABLE}
${CMAKE_CURRENT_SOURCE_DIR}/test_models/gen_wrapper.py
${GEN_SCRIPT}
${TEST_TENSORFLOW_MODELS}
${OUT_DONE_FILE}
JOB_POOL four_jobs
DEPENDS ${TENSORFLOW_ALL_SCRIPTS}
)
endforeach()
add_custom_target(tensorflow_test_models DEPENDS ${OUT_FILES})
install(DIRECTORY ${TEST_TENSORFLOW_MODELS}
DESTINATION tests/${TEST_TENSORFLOW_MODELS_DIRNAME}
COMPONENT tests
EXCLUDE_FROM_ALL)
else()
# Produce warning message at build time as well
add_custom_command(OUTPUT unable_build_tensorflow_models.txt
COMMAND ${CMAKE_COMMAND}
-E cmake_echo_color --red "Warning: Unable to generate tensorflow test models. Running '${TARGET_NAME}' will likely fail"
)
add_custom_target(tensorflow_test_models DEPENDS unable_build_tensorflow_models.txt)
endif()
add_dependencies(${TARGET_NAME} tensorflow_test_models)
add_dependencies(${TARGET_NAME} tensorflow_ngraph_frontend)

View File

@ -0,0 +1,23 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include "basic_api.hpp"
#include "tf_utils.hpp"
using namespace ngraph;
using namespace ngraph::frontend;
using TFBasicTest = FrontEndBasicTest;
static const std::vector<std::string> models{
std::string("2in_2out/2in_2out.pb"),
};
INSTANTIATE_TEST_SUITE_P(TFBasicTest,
FrontEndBasicTest,
::testing::Combine(::testing::Values(TF_FE),
::testing::Values(std::string(TEST_TENSORFLOW_MODELS_DIRNAME)),
::testing::ValuesIn(models)),
FrontEndBasicTest::getTestCaseName);

View File

@ -0,0 +1,23 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include "convert_model.hpp"
#include "tf_utils.hpp"
using namespace ngraph;
using namespace ngraph::frontend;
using TFConvertModelTest = FrontEndConvertModelTest;
static const std::vector<std::string> models{
std::string("2in_2out/2in_2out.pb"),
};
INSTANTIATE_TEST_SUITE_P(TFConvertModelTest,
FrontEndConvertModelTest,
::testing::Combine(::testing::Values(TF_FE),
::testing::Values(std::string(TEST_TENSORFLOW_MODELS_DIRNAME)),
::testing::ValuesIn(models)),
FrontEndConvertModelTest::getTestCaseName);

View File

@ -0,0 +1,39 @@
// Copyright (C) 2018-2021 Intel Corporation
// SPDX-License-Identifier: Apache-2.0
//
#include <frontend_manager/frontend_exceptions.hpp>
#include <frontend_manager/frontend_manager.hpp>
#include "common_test_utils/ngraph_test_utils.hpp"
#include "tf_utils.hpp"
#include "utils.hpp"
using namespace ngraph;
using namespace ngraph::frontend;
TEST(FrontEndConvertModelTest, test_unsupported_op) {
FrontEndManager fem;
FrontEnd::Ptr frontEnd;
InputModel::Ptr inputModel;
ASSERT_NO_THROW(frontEnd = fem.load_by_framework(TF_FE));
ASSERT_NE(frontEnd, nullptr);
auto model_filename = FrontEndTestUtils::make_model_path(std::string(TEST_TENSORFLOW_MODELS_DIRNAME) +
std::string("relu_unsupported/relu_unsupported.pb"));
ASSERT_NO_THROW(inputModel = frontEnd->load(model_filename));
ASSERT_NE(inputModel, nullptr);
std::shared_ptr<ngraph::Function> function;
ASSERT_THROW(function = frontEnd->convert(inputModel), OpConversionFailure);
ASSERT_EQ(function, nullptr);
ASSERT_NO_THROW(function = frontEnd->decode(inputModel));
ASSERT_THROW(frontEnd->convert(function), OpConversionFailure);
ASSERT_NO_THROW(function = frontEnd->convert_partially(inputModel));
ASSERT_THROW(frontEnd->convert(function), OpConversionFailure);
for (auto& node : function->get_ordered_ops()) {
if (node->get_friendly_name() == "relu_0") {
function->replace_node(node, std::make_shared<opset6::Relu>(node->input(0).get_source_output()));
}
}
ASSERT_NO_THROW(frontEnd->convert(function));
}

Some files were not shown because too many files have changed in this diff Show More