[OVC] Fixed error in --output_model param. (#20969)
* Fixed error in get_model_name_from_args(), added tests. * Corrected output_model description. * Corrected test.
This commit is contained in:
@@ -2,7 +2,10 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
@@ -11,7 +14,7 @@ from openvino.runtime import PartialShape, Model
|
||||
from openvino.test_utils import compare_functions
|
||||
from openvino.tools.ovc import ovc
|
||||
|
||||
from common.mo_convert_test_class import CommonMOConvertTest
|
||||
from common import constants
|
||||
from common.tf_layer_test_class import save_to_pb
|
||||
from common.utils.common_utils import shell
|
||||
|
||||
@@ -58,7 +61,12 @@ def create_ref_graph():
|
||||
|
||||
return Model([sigm], [param], "test")
|
||||
|
||||
class TestOVCTool(CommonMOConvertTest):
|
||||
class TestOVCTool(unittest.TestCase):
|
||||
def setUp(self):
|
||||
Path(constants.out_path).mkdir(parents=True, exist_ok=True)
|
||||
test_name = re.sub(r"[^\w_]", "_", unittest.TestCase.id(self))
|
||||
self.tmp_dir = tempfile.TemporaryDirectory(dir=constants.out_path, prefix=f"{test_name}").name
|
||||
|
||||
def create_tf_model(self, tmp_dir):
|
||||
import tensorflow as tf
|
||||
|
||||
@@ -100,58 +108,71 @@ class TestOVCTool(CommonMOConvertTest):
|
||||
return temp_dir + "/test_model", model_ref
|
||||
|
||||
|
||||
def test_ovc_tool(self, ie_device, precision, ir_version, temp_dir, use_new_frontend, use_old_api):
|
||||
def test_ovc_tool(self):
|
||||
from openvino.runtime import Core
|
||||
|
||||
model_path = self.create_tf_model(temp_dir)
|
||||
model_path = self.create_tf_model(self.tmp_dir)
|
||||
|
||||
core = Core()
|
||||
|
||||
# tests for MO cli tool
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_path, "output_model": temp_dir + os.sep + "model1"})
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_path, "output_model": self.tmp_dir + os.sep + "model1"})
|
||||
assert not exit_code
|
||||
|
||||
ov_model = core.read_model(os.path.join(temp_dir, "model1.xml"))
|
||||
ov_model = core.read_model(os.path.join(self.tmp_dir, "model1.xml"))
|
||||
flag, msg = compare_functions(ov_model, create_ref_graph(), False)
|
||||
assert flag, msg
|
||||
|
||||
def test_ovc_tool_output_dir(self, ie_device, precision, ir_version, temp_dir, use_new_frontend, use_old_api):
|
||||
def test_ovc_tool_output_dir(self):
|
||||
from openvino.runtime import Core
|
||||
|
||||
model_path = self.create_tf_model(temp_dir)
|
||||
model_path = self.create_tf_model(self.tmp_dir)
|
||||
|
||||
core = Core()
|
||||
|
||||
# tests for MO cli tool
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_path, "output_model": temp_dir})
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_path, "output_model": self.tmp_dir})
|
||||
assert not exit_code
|
||||
|
||||
ov_model = core.read_model(os.path.join(temp_dir, "model2.xml"))
|
||||
ov_model = core.read_model(os.path.join(self.tmp_dir, "model2.xml"))
|
||||
flag, msg = compare_functions(ov_model, create_ref_graph(), False)
|
||||
assert flag, msg
|
||||
|
||||
def test_ovc_tool_saved_model_dir(self, ie_device, precision, ir_version, temp_dir, use_new_frontend, use_old_api):
|
||||
def test_ovc_tool_saved_model_dir(self):
|
||||
from openvino.runtime import Core
|
||||
core = Core()
|
||||
|
||||
model_dir, ref_model = self.create_tf_saved_model_dir(temp_dir)
|
||||
model_dir, ref_model = self.create_tf_saved_model_dir(self.tmp_dir)
|
||||
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_dir, "output_model": temp_dir})
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_dir, "output_model": self.tmp_dir})
|
||||
assert not exit_code
|
||||
|
||||
ov_model = core.read_model(os.path.join(temp_dir, "test_model.xml"))
|
||||
ov_model = core.read_model(os.path.join(self.tmp_dir, "test_model.xml"))
|
||||
flag, msg = compare_functions(ov_model, ref_model, False)
|
||||
assert flag, msg
|
||||
|
||||
def test_ovc_tool_saved_model_dir_with_sep_at_path_end(self, ie_device, precision, ir_version, temp_dir, use_new_frontend, use_old_api):
|
||||
def test_ovc_tool_saved_model_dir_with_sep_at_path_end(self):
|
||||
from openvino.runtime import Core
|
||||
core = Core()
|
||||
|
||||
model_dir, ref_model = self.create_tf_saved_model_dir(temp_dir)
|
||||
model_dir, ref_model = self.create_tf_saved_model_dir(self.tmp_dir)
|
||||
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_dir + os.sep, "output_model": temp_dir})
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_dir + os.sep, "output_model": self.tmp_dir})
|
||||
assert not exit_code
|
||||
|
||||
ov_model = core.read_model(os.path.join(temp_dir, "test_model.xml"))
|
||||
ov_model = core.read_model(os.path.join(self.tmp_dir, "test_model.xml"))
|
||||
flag, msg = compare_functions(ov_model, ref_model, False)
|
||||
assert flag, msg
|
||||
|
||||
def test_ovc_tool_non_existng_output_dir(self):
|
||||
from openvino.runtime import Core
|
||||
core = Core()
|
||||
|
||||
model_dir, ref_model = self.create_tf_saved_model_dir(self.tmp_dir)
|
||||
|
||||
exit_code, stderr = generate_ir_ovc(coverage=False, **{"input_model": model_dir + os.sep, "output_model": self.tmp_dir + os.sep + "dir" + os.sep})
|
||||
assert not exit_code
|
||||
|
||||
ov_model = core.read_model(os.path.join(self.tmp_dir, "dir", "test_model.xml"))
|
||||
flag, msg = compare_functions(ov_model, ref_model, False)
|
||||
assert flag, msg
|
||||
|
||||
@@ -395,7 +395,9 @@ def get_common_cli_parser(parser: argparse.ArgumentParser = None):
|
||||
|
||||
# Command line tool specific params
|
||||
parser.add_argument('--output_model',
|
||||
help='This parameter is used to name output .xml/.bin files with converted model.')
|
||||
help='This parameter is used to name output .xml/.bin files of converted model. '
|
||||
'Model name or output directory can be passed. If output directory is passed, '
|
||||
'the resulting .xml/.bin files are named by original model name.')
|
||||
parser.add_argument('--compress_to_fp16', type=check_bool, default=True, nargs='?',
|
||||
help='Compress weights in output OpenVINO model to FP16. '
|
||||
'To turn off compression use "--compress_to_fp16=False" command line parameter. '
|
||||
@@ -559,16 +561,21 @@ def get_model_name_from_args(argv: argparse.Namespace):
|
||||
if hasattr(argv, 'output_model') and argv.output_model:
|
||||
model_name = argv.output_model
|
||||
|
||||
if not os.path.isdir(argv.output_model):
|
||||
if not os.path.isdir(argv.output_model) and not argv.output_model.endswith(os.sep):
|
||||
# In this branch we assume that model name is set in 'output_model'.
|
||||
if not model_name.endswith('.xml'):
|
||||
model_name += '.xml'
|
||||
# Logic of creating and checking directory is covered in save_model() method.
|
||||
return model_name
|
||||
else:
|
||||
if not os.access(argv.output_model, os.W_OK):
|
||||
# In this branch 'output_model' has directory without name of model.
|
||||
# The directory may not exist.
|
||||
if os.path.isdir(argv.output_model) and not os.access(argv.output_model, os.W_OK):
|
||||
# If the provided path is existing directory, but not writable, then raise error
|
||||
raise Error('The directory "{}" is not writable'.format(argv.output_model))
|
||||
output_dir = argv.output_model
|
||||
|
||||
input_model = argv.input_model
|
||||
input_model = os.path.abspath(argv.input_model)
|
||||
if isinstance(input_model, (tuple, list)) and len(input_model) > 0:
|
||||
input_model = input_model[0]
|
||||
|
||||
@@ -583,8 +590,8 @@ def get_model_name_from_args(argv: argparse.Namespace):
|
||||
input_model_name = os.path.splitext(input_model_name)[0]
|
||||
|
||||
# if no valid name exists in input path set name to 'model'
|
||||
if input_model_name == '' or input_model_name == '.':
|
||||
input_model_name = "model"
|
||||
if input_model_name == '':
|
||||
raise Exception("Could not derive model name from input model. Please provide 'output_model' parameter.")
|
||||
|
||||
# add .xml extension
|
||||
return os.path.join(output_dir, input_model_name + ".xml")
|
||||
|
||||
@@ -2,24 +2,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import argparse
|
||||
import numpy
|
||||
import os
|
||||
import shutil
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
import numpy as np
|
||||
import openvino.runtime as ov
|
||||
from openvino.runtime import PartialShape
|
||||
|
||||
from openvino.tools.ovc.cli_parser import _InputCutInfo
|
||||
from openvino.tools.ovc.cli_parser import input_to_input_cut_info, \
|
||||
readable_file_or_object, get_all_cli_parser, get_mo_convert_params, parse_inputs
|
||||
readable_file_or_object, get_all_cli_parser, get_mo_convert_params, parse_inputs, get_model_name_from_args
|
||||
from openvino.tools.ovc.convert_impl import pack_params_to_args_namespace, arguments_post_parsing, args_to_argv
|
||||
from openvino.tools.ovc.error import Error
|
||||
from unit_tests.ovc.unit_test_with_mocked_telemetry import UnitTestWithMockedTelemetry
|
||||
from openvino.runtime import PartialShape, Dimension, Layout
|
||||
from openvino.tools.ovc.cli_parser import _InputCutInfo
|
||||
import openvino.runtime as ov
|
||||
|
||||
|
||||
class TestShapesParsing(UnitTestWithMockedTelemetry):
|
||||
@@ -518,3 +514,76 @@ class TestConvertModelParamsParsing(unittest.TestCase):
|
||||
assert param_name not in cli_parser._option_string_actions
|
||||
else:
|
||||
assert param_name in cli_parser._option_string_actions
|
||||
|
||||
|
||||
|
||||
class GetModelNameTest(unittest.TestCase):
|
||||
def test_case1(self):
|
||||
current_dir = os.getcwd()
|
||||
dir = os.path.basename(current_dir)
|
||||
argv = argparse.Namespace(input_model="."+ os.sep)
|
||||
assert get_model_name_from_args(argv) == current_dir + os.sep + dir + ".xml"
|
||||
|
||||
def test_case2(self):
|
||||
current_dir = os.getcwd()
|
||||
argv = argparse.Namespace(input_model="."+ os.sep +"test_model")
|
||||
assert get_model_name_from_args(argv) == current_dir + os.sep + "test_model.xml"
|
||||
|
||||
|
||||
def test_case3(self):
|
||||
current_dir = os.getcwd()
|
||||
argv = argparse.Namespace(input_model="."+ os.sep +"test_model.pb")
|
||||
assert get_model_name_from_args(argv) == current_dir + os.sep + "test_model.xml"
|
||||
|
||||
|
||||
def test_case4(self):
|
||||
current_dir = os.getcwd()
|
||||
dir = os.path.basename(current_dir)
|
||||
argv = argparse.Namespace(input_model="."+ os.sep,
|
||||
output_model="."+ os.sep)
|
||||
assert get_model_name_from_args(argv) == "." + os.sep + dir + ".xml"
|
||||
|
||||
def test_case5(self):
|
||||
argv = argparse.Namespace(input_model="test_model.pb",
|
||||
output_model="."+ os.sep)
|
||||
assert get_model_name_from_args(argv) == "." + os.sep + "test_model.xml"
|
||||
|
||||
|
||||
def test_case6(self):
|
||||
argv = argparse.Namespace(input_model="test_model",
|
||||
output_model="."+ os.sep)
|
||||
assert get_model_name_from_args(argv) == "." + os.sep + "test_model.xml"
|
||||
|
||||
|
||||
|
||||
def test_case7(self):
|
||||
argv = argparse.Namespace(input_model="test_dir" + os.sep,
|
||||
output_model="."+ os.sep)
|
||||
assert get_model_name_from_args(argv) == "." + os.sep + "test_dir.xml"
|
||||
|
||||
def test_case8(self):
|
||||
argv = argparse.Namespace(input_model="test_model.pb",
|
||||
output_model="new_model")
|
||||
assert get_model_name_from_args(argv) == "new_model.xml"
|
||||
|
||||
def test_case9(self):
|
||||
argv = argparse.Namespace(input_model="test_model",
|
||||
output_model="new_dir" + os.sep)
|
||||
assert get_model_name_from_args(argv) == "new_dir" + os.sep + "test_model.xml"
|
||||
|
||||
def test_case10(self):
|
||||
argv = argparse.Namespace(input_model="test_dir" + os.sep,
|
||||
output_model="new_model.xml")
|
||||
assert get_model_name_from_args(argv) == "new_model.xml"
|
||||
|
||||
def test_case11(self):
|
||||
argv = argparse.Namespace(input_model="/",
|
||||
output_model="new_model")
|
||||
assert get_model_name_from_args(argv) == "new_model.xml"
|
||||
|
||||
|
||||
def test_negative(self):
|
||||
|
||||
argv = argparse.Namespace(input_model="/",)
|
||||
with self.assertRaisesRegex(Exception, ".*Could not derive model name from input model. Please provide 'output_model' parameter.*"):
|
||||
get_model_name_from_args(argv)
|
||||
Reference in New Issue
Block a user