From 356386602fa36c1471ac18682f521fe47d768498 Mon Sep 17 00:00:00 2001 From: danielsongmicrosoft Date: Sun, 27 Sep 2026 22:24:48 -0700 Subject: [PATCH] Check CPU fallback assignments in nested subgraphs --- onnxruntime/core/session/inference_session.cc | 15 +- .../core/session/inference_session_utils.cc | 17 ++ .../core/session/inference_session_utils.h | 2 + .../test/framework/inference_session_test.cc | 202 ++++++++++++++++++ 4 files changed, 222 insertions(+), 14 deletions(-) diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index 16202678e0153..23b07d984a221 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -3132,19 +3132,6 @@ common::Status InferenceSession::Initialize() { // If the user disabled fallback, but also explicitly added the CPU EP to the session, return an error status. // If the user disabled fallback and any graph node is assigned to the CPU EP, return an error status. if (disable_cpu_ep_fallback) { - // Returns true if any graph nodes have been assigned to the CPU EP. - auto are_nodes_assigned_to_cpu_ep = [](const Graph& graph) -> bool { - for (const auto& node : graph.Nodes()) { - const auto& node_provider = node.GetExecutionProviderType(); - - if (node_provider.empty() || node_provider == onnxruntime::kCpuExecutionProvider) { - return true; - } - } - - return false; - }; - if (!execution_providers_.GetCpuProviderWasImplicitlyAdded()) { const char* err_msg = "Conflicting session configuration: explicitly added the CPU EP to the " @@ -3152,7 +3139,7 @@ common::Status InferenceSession::Initialize() { LOGS(*session_logger_, ERROR) << err_msg; ORT_RETURN_IF_ERROR_SESSIONID_(ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, err_msg)); - } else if (are_nodes_assigned_to_cpu_ep(graph)) { + } else if (inference_session_utils::AreAnyNodesAssignedToCpuEp(graph)) { const char* err_msg = "This session contains graph nodes that are assigned to the default CPU EP, " "but fallback to CPU EP has been explicitly disabled by the user."; diff --git a/onnxruntime/core/session/inference_session_utils.cc b/onnxruntime/core/session/inference_session_utils.cc index 579a4fc11d3ee..f86435b9ecf14 100644 --- a/onnxruntime/core/session/inference_session_utils.cc +++ b/onnxruntime/core/session/inference_session_utils.cc @@ -7,6 +7,23 @@ namespace onnxruntime { +bool inference_session_utils::AreAnyNodesAssignedToCpuEp(const Graph& graph) { + for (const auto& node : graph.Nodes()) { + const auto& node_provider = node.GetExecutionProviderType(); + if (node_provider.empty() || node_provider == kCpuExecutionProvider) { + return true; + } + + for (const gsl::not_null& subgraph : node.GetSubgraphs()) { + if (AreAnyNodesAssignedToCpuEp(*subgraph)) { + return true; + } + } + } + + return false; +} + //--------------------- //--- local helpers --- //--------------------- diff --git a/onnxruntime/core/session/inference_session_utils.h b/onnxruntime/core/session/inference_session_utils.h index f297d928f8a0d..592859dd1e25d 100644 --- a/onnxruntime/core/session/inference_session_utils.h +++ b/onnxruntime/core/session/inference_session_utils.h @@ -28,6 +28,8 @@ namespace inference_session_utils { static constexpr const char* kOrtLoadConfigFromModelEnvVar = "ORT_LOAD_CONFIG_FROM_MODEL"; #if !defined(ORT_MINIMAL_BUILD) +bool AreAnyNodesAssignedToCpuEp(const Graph& graph); + // // Code to parse json session config from onnx model file // diff --git a/onnxruntime/test/framework/inference_session_test.cc b/onnxruntime/test/framework/inference_session_test.cc index 3229523b9f798..cd1d3a7f8e771 100644 --- a/onnxruntime/test/framework/inference_session_test.cc +++ b/onnxruntime/test/framework/inference_session_test.cc @@ -132,6 +132,32 @@ ONNX_OPERATOR_KERNEL_EX(FuseAdd, // .TypeConstraint("T", DataTypeImpl::GetTensorType()), FuseAdd); +class DisableCpuFallbackKernel : public OpKernel { + public: + explicit DisableCpuFallbackKernel(const OpKernelInfo& info) : OpKernel(info) { + } + + Status Compute(OpKernelContext* /*context*/) const override { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Test kernel should not be executed"); + } +}; + +constexpr const char* kDisableCpuFallbackExecutionProvider = "DisableCpuFallbackExecutionProvider"; +class ONNX_OPERATOR_KERNEL_CLASS_NAME(kDisableCpuFallbackExecutionProvider, kOnnxDomain, 13, If); +class ONNX_OPERATOR_KERNEL_CLASS_NAME(kDisableCpuFallbackExecutionProvider, kOnnxDomain, 13, Abs); +ONNX_OPERATOR_KERNEL_EX(If, + kOnnxDomain, + 13, + kDisableCpuFallbackExecutionProvider, + KernelDefBuilder(), + DisableCpuFallbackKernel); +ONNX_OPERATOR_KERNEL_EX(Abs, + kOnnxDomain, + 13, + kDisableCpuFallbackExecutionProvider, + KernelDefBuilder(), + DisableCpuFallbackKernel); + Status RegisterOperatorKernels(KernelRegistry& kernel_registry) { return kernel_registry.Register( BuildKernelCreateInfo()); @@ -143,6 +169,26 @@ KernelRegistryAndStatus GetFusedKernelRegistry() { return ret; } +Status RegisterDisableCpuFallbackTestKernels(KernelRegistry& kernel_registry) { + ORT_RETURN_IF_ERROR( + kernel_registry.Register(BuildKernelCreateInfo< + ONNX_OPERATOR_KERNEL_CLASS_NAME(kDisableCpuFallbackExecutionProvider, + kOnnxDomain, + 13, + If)>())); + return kernel_registry.Register(BuildKernelCreateInfo< + ONNX_OPERATOR_KERNEL_CLASS_NAME(kDisableCpuFallbackExecutionProvider, + kOnnxDomain, + 13, + Abs)>()); +} + +KernelRegistryAndStatus GetDisableCpuFallbackTestKernelRegistry() { + KernelRegistryAndStatus ret; + ret.st = RegisterDisableCpuFallbackTestKernels(*ret.kernel_registry); + return ret; +} + class FuseExecutionProvider : public IExecutionProvider { public: explicit FuseExecutionProvider() : IExecutionProvider{kFuseExecutionProvider} { @@ -195,6 +241,45 @@ class FuseExecutionProvider : public IExecutionProvider { } }; +class DisableCpuFallbackTestExecutionProvider : public IExecutionProvider { + public: + explicit DisableCpuFallbackTestExecutionProvider(bool claim_leaf_ops) + : IExecutionProvider{kDisableCpuFallbackExecutionProvider}, + claim_leaf_ops_{claim_leaf_ops} { + } + + std::vector> GetCapability( + const onnxruntime::GraphViewer& graph, + const IKernelLookup& /*kernel_lookup*/, + const GraphOptimizerRegistry& /* graph_optimizer_registry */, + IResourceAccountant* /* resource_accountant */) const override { + std::vector> result; + + for (const auto& node : graph.Nodes()) { + const bool should_claim = node.OpType() == "If" || + (claim_leaf_ops_ && node.OpType() == "Abs"); + if (!should_claim) { + continue; + } + + auto sub_graph = std::make_unique(); + sub_graph->nodes.push_back(node.Index()); + result.push_back(std::make_unique(std::move(sub_graph))); + } + + return result; + } + + std::shared_ptr GetKernelRegistry() const override { + static KernelRegistryAndStatus k = GetDisableCpuFallbackTestKernelRegistry(); + ORT_THROW_IF_ERROR(k.st); + return k.kernel_registry; + } + + private: + bool claim_leaf_ops_; +}; + namespace test { static constexpr const ORTCHAR_T* MODEL_URI = ORT_TSTR("testdata/mul_1.onnx"); static constexpr const ORTCHAR_T* MODEL_URI_NO_OPSET = ORT_TSTR("testdata/mul_1.noopset.onnx"); @@ -1901,6 +1986,123 @@ TEST(InferenceSessionTests, TestOptionalInputs) { } } +static void CreateNestedIfModel(const PathString& model_file_name) { + ONNX_NAMESPACE::TypeProto bool_tensor; + bool_tensor.mutable_tensor_type()->set_elem_type(ONNX_NAMESPACE::TensorProto_DataType_BOOL); + bool_tensor.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(1); + + ONNX_NAMESPACE::TypeProto float_tensor; + float_tensor.mutable_tensor_type()->set_elem_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT); + float_tensor.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(1); + + ONNX_NAMESPACE::GraphProto leaf_graph; + { + onnxruntime::Model model("nested_if_leaf_graph", false, ModelMetaData(), PathString(), + IOnnxRuntimeOpSchemaRegistryList(), {{kOnnxDomain, 13}}, {}, + DefaultLoggingManager().DefaultLogger()); + auto& graph = model.MainGraph(); + + auto& branch_data = graph.GetOrCreateNodeArg("branch_data", &float_tensor); + graph.AddOuterScopeNodeArg("branch_data"); + auto& branch_output = graph.GetOrCreateNodeArg("branch_output", &float_tensor); + + graph.AddNode("branch_abs", "Abs", "Abs node in nested branch", {&branch_data}, {&branch_output}); + graph.SetOutputs({&branch_output}); + + ASSERT_STATUS_OK(graph.Resolve()); + leaf_graph = graph.ToGraphProto(); + } + + ONNX_NAMESPACE::GraphProto middle_graph; + { + onnxruntime::Model model("nested_if_middle_graph", false, ModelMetaData(), PathString(), + IOnnxRuntimeOpSchemaRegistryList(), {{kOnnxDomain, 13}}, {}, + DefaultLoggingManager().DefaultLogger()); + auto& graph = model.MainGraph(); + + auto& inner_cond = graph.GetOrCreateNodeArg("inner_cond", &bool_tensor); + auto& branch_data = graph.GetOrCreateNodeArg("branch_data", &float_tensor); + ORT_UNUSED_PARAMETER(inner_cond); + ORT_UNUSED_PARAMETER(branch_data); + graph.AddOuterScopeNodeArg("inner_cond"); + graph.AddOuterScopeNodeArg("branch_data"); + + auto& middle_output = graph.GetOrCreateNodeArg("middle_output", &float_tensor); + auto& inner_if = graph.AddNode("middle_if", "If", "Inner If node", {&inner_cond}, {&middle_output}); + inner_if.AddAttribute("then_branch", leaf_graph); + inner_if.AddAttribute("else_branch", leaf_graph); + + graph.SetOutputs({&middle_output}); + + ASSERT_STATUS_OK(graph.Resolve()); + middle_graph = graph.ToGraphProto(); + } + + onnxruntime::Model model("nested_if_main_graph", false, ModelMetaData(), PathString(), + IOnnxRuntimeOpSchemaRegistryList(), {{kOnnxDomain, 13}}, {}, + DefaultLoggingManager().DefaultLogger()); + auto& graph = model.MainGraph(); + + auto& outer_cond = graph.GetOrCreateNodeArg("outer_cond", &bool_tensor); + auto& inner_cond = graph.GetOrCreateNodeArg("inner_cond", &bool_tensor); + auto& branch_data = graph.GetOrCreateNodeArg("branch_data", &float_tensor); + ORT_UNUSED_PARAMETER(inner_cond); + ORT_UNUSED_PARAMETER(branch_data); + auto& output = graph.GetOrCreateNodeArg("output", &float_tensor); + + auto& outer_if = graph.AddNode("outer_if", "If", "Outer If node", {&outer_cond}, {&output}); + outer_if.AddAttribute("then_branch", middle_graph); + outer_if.AddAttribute("else_branch", middle_graph); + + graph.SetInputs({&outer_cond, &inner_cond, &branch_data}); + graph.SetOutputs({&output}); + + ASSERT_STATUS_OK(graph.Resolve()); + ASSERT_STATUS_OK(onnxruntime::Model::Save(model, model_file_name)); +} + +TEST(InferenceSessionTests, DisableCpuEpFallbackRejectsCpuNodesInNestedSubgraphs) { + const PathString model_file_name = ORT_TSTR("disable_cpu_ep_fallback_nested_if.onnx"); + CreateNestedIfModel(model_file_name); + + SessionOptions so; + so.session_logid = "InferenceSessionTests.DisableCpuEpFallbackRejectsCpuNodesInNestedSubgraphs"; + ASSERT_STATUS_OK(so.config_options.AddConfigEntry(kOrtSessionOptionsDisableCPUEPFallback, "1")); + + InferenceSession session_object{so, GetEnvironment()}; + ASSERT_STATUS_OK(session_object.RegisterExecutionProvider( + std::make_unique(false))); + ASSERT_STATUS_OK(session_object.Load(model_file_name)); + ASSERT_STATUS_NOT_OK_AND_HAS_SUBSTR(session_object.Initialize(), + "fallback to CPU EP has been explicitly disabled"); +} + +TEST(InferenceSessionTests, DisableCpuEpFallbackAllowsFullyAssignedNestedSubgraphs) { + const PathString model_file_name = ORT_TSTR("disable_cpu_ep_fallback_nested_if_fully_assigned.onnx"); + CreateNestedIfModel(model_file_name); + + SessionOptions so; + so.session_logid = "InferenceSessionTests.DisableCpuEpFallbackAllowsFullyAssignedNestedSubgraphs"; + ASSERT_STATUS_OK(so.config_options.AddConfigEntry(kOrtSessionOptionsDisableCPUEPFallback, "1")); + + InferenceSessionWrapper session_object{so, GetEnvironment()}; + ASSERT_STATUS_OK(session_object.Load(model_file_name)); + + std::function assign_to_test_ep = [&](Graph& graph) { + for (auto& node : graph.Nodes()) { + node.SetExecutionProviderType(kDisableCpuFallbackExecutionProvider); + for (const auto& [attribute_name, subgraph] : node.GetAttributeNameToMutableSubgraphMap()) { + ORT_UNUSED_PARAMETER(attribute_name); + assign_to_test_ep(*subgraph); + } + } + }; + auto& graph = session_object.GetMutableGraph(); + assign_to_test_ep(graph); + + EXPECT_FALSE(inference_session_utils::AreAnyNodesAssignedToCpuEp(graph)); +} + static void CreateFuseOpModel(const PathString& model_file_name) { onnxruntime::Model model("graph_1", false, ModelMetaData(), PathString(), IOnnxRuntimeOpSchemaRegistryList(), {{kOnnxDomain, 12}}, {}, DefaultLoggingManager().DefaultLogger());