| 787 | } |
| 788 | |
| 789 | tensorflow::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 | } |
nothing calls this directly
no test coverage detected