diff --git a/mindspore/ccsrc/pipeline/jit/static_analysis/program_specialize.cc b/mindspore/ccsrc/pipeline/jit/static_analysis/program_specialize.cc index 2536a002e0..4fb8b276cc 100644 --- a/mindspore/ccsrc/pipeline/jit/static_analysis/program_specialize.cc +++ b/mindspore/ccsrc/pipeline/jit/static_analysis/program_specialize.cc @@ -51,6 +51,19 @@ bool IsVisible(FuncGraphPtr fg, const FuncGraphPtr &parent) { } return fg == parent; } + +bool CheckAbstractTensor(const AbstractBasePtr &abs_base) { + if (abs_base->isa()) { + return true; + } else if (abs_base->isa()) { + const auto &abs_seq = abs_base->cast(); + MS_EXCEPTION_IF_NULL(abs_seq); + const auto &elements = abs_seq->elements(); + return std::all_of(elements.cbegin(), elements.cend(), [](const auto &v) { return CheckAbstractTensor(v); }); + } else { + return false; + } +} } // namespace FuncGraphPtr ProgramSpecializer::Run(const FuncGraphPtr &fg, const AnalysisContextPtr &context) { @@ -575,7 +588,7 @@ std::pair FuncGraphSpecializer::BuildFromB return std::make_pair(joined_argvals, joined_eval_result->abstract()); } else { bool all_args_tensor = std::all_of(broaded_argvals.cbegin(), broaded_argvals.cend(), - [](const AbstractBasePtr &v) { return v->isa(); }); + [](const AbstractBasePtr &v) { return CheckAbstractTensor(v); }); if (all_args_tensor) { ConfigPtrList args_conf_list; (void)std::transform(broaded_argvals.cbegin(), broaded_argvals.cend(), std ::back_inserter(args_conf_list),