From da36633c08996138f1fd94df5ce1d3a38b447faa Mon Sep 17 00:00:00 2001 From: Maxim Vafin Date: Fri, 4 Aug 2023 09:36:09 +0200 Subject: [PATCH] [PT FE] Raise graceful exception when model has incorrect type (#18976) * [PT FE] Raise graceful exception when model has incorrect type * Improve message --- src/frontends/pytorch/src/frontend.cpp | 3 +++ .../mo_python_api_tests/test_mo_convert_pytorch.py | 8 ++++++++ 2 files changed, 11 insertions(+) diff --git a/src/frontends/pytorch/src/frontend.cpp b/src/frontends/pytorch/src/frontend.cpp index a519a34eed7..2ddf81670c0 100644 --- a/src/frontends/pytorch/src/frontend.cpp +++ b/src/frontends/pytorch/src/frontend.cpp @@ -254,6 +254,9 @@ ov::frontend::InputModel::Ptr FrontEnd::load_impl(const std::vector& va "PyTorch Frontend supports exactly one parameter in model representation, got ", std::to_string(variants.size()), " instead."); + FRONT_END_GENERAL_CHECK(variants[0].is>(), + "PyTorch Frontend doesn't support provided model type. Please provide supported model " + "object using Python API."); auto decoder = variants[0].as>(); auto tdecoder = std::dynamic_pointer_cast(decoder); FRONT_END_GENERAL_CHECK(tdecoder, "Couldn't cast ov::Any to TorchDecoder"); diff --git a/tests/layer_tests/mo_python_api_tests/test_mo_convert_pytorch.py b/tests/layer_tests/mo_python_api_tests/test_mo_convert_pytorch.py index 9b0d37c2bf9..5e6cdc765c9 100644 --- a/tests/layer_tests/mo_python_api_tests/test_mo_convert_pytorch.py +++ b/tests/layer_tests/mo_python_api_tests/test_mo_convert_pytorch.py @@ -1106,3 +1106,11 @@ class ConvertRaises(unittest.TestCase): with self.assertRaisesRegex(Exception, ".*Conversion is failed for: aten::relu.*"): convert_model(pt_model, input=(inp_shapes, np.float32), extensions=[ ConversionExtension("aten::relu", relu_bad)]) + + def test_failed_extension(self): + import tempfile + from openvino.tools.mo import convert_model + + with self.assertRaisesRegex(Exception, ".*PyTorch Frontend doesn't support provided model type.*"): + with tempfile.NamedTemporaryFile() as tmpfile: + convert_model(tmpfile.name, framework="pytorch")