MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / ConvertConvOperator

Function ConvertConvOperator

tensorflow/lite/toco/import_tensorflow.cc:789–868  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

787}
788
789tensorflow::Status ConvertConvOperator(
790 const NodeDef& node, const TensorFlowImportFlags& tf_import_flags,
791 const ModelFlags& model_flags, Model* model) {
792 CHECK_EQ(node.op(), "Conv2D");
793 TF_RETURN_IF_ERROR(CheckInputsCount(node, tf_import_flags, 2));
794
795 // We only support NHWC, which is the default data_format.
796 // So if data_format is not defined, we're all good.
797 TF_RETURN_IF_ERROR(CheckOptionalAttr(node, "data_format", "NHWC"));
798 TF_RETURN_IF_ERROR(CheckOptionalAttr(node, "T", DT_FLOAT));
799
800 const auto& input_name = node.input(0);
801 const auto& weights_name = node.input(1);
802 const auto& reordered_weights_name =
803 AvailableArrayName(*model, weights_name + "_reordered");
804 // Check if a ReorderAxesOperator was already created for these weights
805 // (that happens when multiple layers share the same weights).
806 const Operator* existing_reorder =
807 GetOpWithOutput(*model, reordered_weights_name);
808 if (existing_reorder) {
809 // Check that it is safe to rely on the _reordered naming of the output
810 // array!
811 CHECK(existing_reorder->type == OperatorType::kReorderAxes);
812 } else {
813 // Create a new ReorderAxesOperator
814 auto* reorder = new ReorderAxesOperator;
815 reorder->inputs = {weights_name};
816 reorder->outputs = {reordered_weights_name};
817 reorder->input_axes_order = AxesOrder::kHWIO;
818 reorder->output_axes_order = AxesOrder::kOHWI;
819 model->operators.emplace_back(reorder);
820 }
821 if (!HasAttr(node, "strides")) {
822 return tensorflow::errors::InvalidArgument("Missing attribute 'strides'");
823 }
824 const auto& strides = GetListAttr(node, "strides");
825 TF_RETURN_IF_ERROR(ExpectValue(strides.i_size(), 4, "number of strides"));
826 TF_RETURN_IF_ERROR(ExpectValue(strides.i(0), 1, "strides(0)"));
827 TF_RETURN_IF_ERROR(ExpectValue(strides.i(3), 1, "strides(3)"));
828 int dilation_height_factor;
829 int dilation_width_factor;
830 if (HasAttr(node, "dilations")) {
831 const auto& dilations = GetListAttr(node, "dilations");
832 TF_RETURN_IF_ERROR(
833 ExpectValue(dilations.i_size(), 4, "number of dilations"));
834 if (dilations.i(0) != 1 || dilations.i(3) != 1) {
835 return tensorflow::errors::InvalidArgument(absl::StrCat(
836 "Can only import Conv ops with dilation along the height "
837 "(1st) or width (2nd) axis. TensorFlow op \"",
838 node.name(), "\" had dilations:[ ", dilations.i(0), ", ",
839 dilations.i(1), ", ", dilations.i(2), ", ", dilations.i(3), "]."));
840 }
841 dilation_height_factor = dilations.i(1);
842 dilation_width_factor = dilations.i(2);
843 } else {
844 dilation_height_factor = 1;
845 dilation_width_factor = 1;
846 }

Callers

nothing calls this directly

Calls 14

CheckInputsCountFunction · 0.85
CheckOptionalAttrFunction · 0.85
AvailableArrayNameFunction · 0.85
GetOpWithOutputFunction · 0.85
HasAttrFunction · 0.85
InvalidArgumentFunction · 0.85
ExpectValueFunction · 0.85
GetStringAttrFunction · 0.85
nameMethod · 0.65
StrCatFunction · 0.50
opMethod · 0.45
inputMethod · 0.45

Tested by

no test coverage detected