Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 1 addition & 14 deletions onnxruntime/core/session/inference_session.cc
Original file line number Diff line number Diff line change
Expand Up @@ -3132,27 +3132,14 @@ 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 "
"session, but also disabled fallback to the CPU EP via session configuration options.";

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.";
Expand Down
17 changes: 17 additions & 0 deletions onnxruntime/core/session/inference_session_utils.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<const Graph*>& subgraph : node.GetSubgraphs()) {
if (AreAnyNodesAssignedToCpuEp(*subgraph)) {
return true;
}
}
}

return false;
}

//---------------------
//--- local helpers ---
//---------------------
Expand Down
2 changes: 2 additions & 0 deletions onnxruntime/core/session/inference_session_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -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
//
Expand Down
202 changes: 202 additions & 0 deletions onnxruntime/test/framework/inference_session_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,32 @@ ONNX_OPERATOR_KERNEL_EX(FuseAdd,
// .TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
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<ONNX_OPERATOR_KERNEL_CLASS_NAME(kFuseExecutionProvider, kFuseTest, 1, FuseAdd)>());
Expand All @@ -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} {
Expand Down Expand Up @@ -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<std::unique_ptr<ComputeCapability>> GetCapability(
const onnxruntime::GraphViewer& graph,
const IKernelLookup& /*kernel_lookup*/,
const GraphOptimizerRegistry& /* graph_optimizer_registry */,
IResourceAccountant* /* resource_accountant */) const override {
std::vector<std::unique_ptr<ComputeCapability>> 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<IndexedSubGraph>();
sub_graph->nodes.push_back(node.Index());
result.push_back(std::make_unique<ComputeCapability>(std::move(sub_graph)));
}

return result;
}

std::shared_ptr<KernelRegistry> 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");
Expand Down Expand Up @@ -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<DisableCpuFallbackTestExecutionProvider>(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<void(Graph&)> 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());
Expand Down
Loading