diff --git a/mindspore/ccsrc/utils/convert_utils_py.cc b/mindspore/ccsrc/utils/convert_utils_py.cc index 88a18096e75..d50e0becce1 100644 --- a/mindspore/ccsrc/utils/convert_utils_py.cc +++ b/mindspore/ccsrc/utils/convert_utils_py.cc @@ -238,6 +238,8 @@ static ValueNameToConverterVector value_name_to_converter = { {AnyValue::kTypeId, [](const ValuePtr &) -> py::object { return py::none(); }}, // FuncGraph {FuncGraph::kTypeId, [](const ValuePtr &) -> py::object { return py::none(); }}, + // Primitive + {Primitive::kTypeId, [](const ValuePtr &) -> py::object { return py::none(); }}, // Monad {Monad::kTypeId, [](const ValuePtr &) -> py::object { return py::none(); }}, // Ellipsis diff --git a/tests/st/fallback/test_graph_fallback.py b/tests/st/fallback/test_graph_fallback.py index 2b5afe4b9cc..81c1f7fdb98 100644 --- a/tests/st/fallback/test_graph_fallback.py +++ b/tests/st/fallback/test_graph_fallback.py @@ -1,4 +1,4 @@ -# Copyright 2021 Huawei Technologies Co., Ltd +# Copyright 2021-2022 Huawei Technologies Co., Ltd # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -17,8 +17,9 @@ import pytest import numpy as np import mindspore.nn as nn -from mindspore import Tensor, ms_function, context import mindspore.common.dtype as mstype +from mindspore import Tensor, ms_function, context +from mindspore.ops import Primitive context.set_context(mode=context.GRAPH_MODE) @@ -257,3 +258,28 @@ def test_fallback_tensor_array_astype(): me_x = Tensor([1.1, -2.1]).astype("float32") return me_x print(foo()) + + +@pytest.mark.level0 +@pytest.mark.platform_x86_gpu_training +@pytest.mark.platform_arm_ascend_training +@pytest.mark.platform_x86_ascend_training +@pytest.mark.env_onecard +def test_fallback_tuple_with_mindspore_function(): + """ + Feature: JIT Fallback + Description: Test fallback when local input has tuple with mindspore function type, such as Cell, Primitive. + Expectation: No exception. + """ + def test_isinstance(a, base_type): + mro = type(a).mro() + for i in base_type: + if i in mro: + return True + return False + + @ms_function + def foo(): + return test_isinstance(np.array(1), (np.ndarray, nn.Cell, Primitive)) + + assert foo()