diff --git a/onnxruntime/core/framework/transformer_memcpy.cc b/onnxruntime/core/framework/transformer_memcpy.cc index d6fcdc7f2a..3c66849a4d 100644 --- a/onnxruntime/core/framework/transformer_memcpy.cc +++ b/onnxruntime/core/framework/transformer_memcpy.cc @@ -116,12 +116,23 @@ bool TransformerMemcpyImpl::ModifyGraph(const KernelRegistryManager& kernel_regi // for initializers shared by different providers, create dups ProcessInitializers(); + for (auto arg : graph_.GetInputs()) + BuildDefsMapping(arg, kernel_registries); + for (auto arg : non_provider_input_defs_) BuildDefsMapping(arg, kernel_registries); for (auto arg : non_provider_output_defs_) BuildDefsMapping(arg, kernel_registries); + for (auto arg : graph_.GetInputs()) + // For inputs we need to create a copy node only when the input is connected to both provider + // and non-provider nodes. Otherwise utils::CopyInputsAcrossDevices() will do the job. + if (provider_input_defs_.count(arg) && non_provider_input_defs_.count(arg)) { + AddCopyNode(const_cast(arg), true); + modified = true; + } + for (auto arg : non_provider_output_defs_) if (provider_input_defs_.count(arg)) { AddCopyNode(arg, true); diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index 4c5e028e8a..b8b9328e5a 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -327,8 +327,8 @@ class InferenceSession::Impl { if (!execution_providers_.Get(onnxruntime::kCpuExecutionProvider)) { LOGS(*session_logger_, INFO) << "Adding default CPU execution provider."; CPUExecutionProviderInfo epi{session_options_.enable_cpu_mem_arena}; - execution_providers_.Add(onnxruntime::kCpuExecutionProvider, - std::make_unique(epi)); + ORT_RETURN_IF_ERROR(execution_providers_.Add(onnxruntime::kCpuExecutionProvider, + std::make_unique(epi))); } onnxruntime::Graph& graph = model_->MainGraph(); diff --git a/onnxruntime/test/framework/inference_session_test.cc b/onnxruntime/test/framework/inference_session_test.cc index 3a776de861..e57c9fa9a5 100644 --- a/onnxruntime/test/framework/inference_session_test.cc +++ b/onnxruntime/test/framework/inference_session_test.cc @@ -1000,17 +1000,17 @@ TEST(ExecutionProviderTest, FunctionTest) { RunOptions run_options; run_options.run_tag = so.session_logid; + CPUExecutionProviderInfo epi; + auto testCPUExecutionProvider = std::make_unique<::onnxruntime::CPUExecutionProvider>(epi); + std::vector dims_mul_x = {3, 2}; std::vector values_mul_x = {1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f}; MLValue ml_value_x; - CreateMLValue(TestCPUExecutionProvider()->GetAllocator(0, OrtMemTypeDefault), dims_mul_x, values_mul_x, - &ml_value_x); + CreateMLValue(testCPUExecutionProvider->GetAllocator(0, OrtMemTypeDefault), dims_mul_x, values_mul_x, &ml_value_x); MLValue ml_value_y; - CreateMLValue(TestCPUExecutionProvider()->GetAllocator(0, OrtMemTypeDefault), dims_mul_x, values_mul_x, - &ml_value_y); + CreateMLValue(testCPUExecutionProvider->GetAllocator(0, OrtMemTypeDefault), dims_mul_x, values_mul_x, &ml_value_y); MLValue ml_value_z; - CreateMLValue(TestCPUExecutionProvider()->GetAllocator(0, OrtMemTypeDefault), dims_mul_x, values_mul_x, - &ml_value_z); + CreateMLValue(testCPUExecutionProvider->GetAllocator(0, OrtMemTypeDefault), dims_mul_x, values_mul_x, &ml_value_z); NameMLValMap feeds; feeds.insert(std::make_pair("X", ml_value_x)); feeds.insert(std::make_pair("Y", ml_value_y)); @@ -1031,7 +1031,8 @@ TEST(ExecutionProviderTest, FunctionTest) { VerifyOutputs(fetches, expected_dims_mul_m, expected_values_mul_m); InferenceSession session_object_2{so}; - session_object_2.RegisterExecutionProvider(std::make_unique()); + session_object_2.RegisterExecutionProvider(std::move(testCPUExecutionProvider)); + session_object_2.RegisterExecutionProvider(std::make_unique<::onnxruntime::FuseExecutionProvider>()); status = session_object_2.Load(model_file_name); ASSERT_TRUE(status.IsOK()); status = session_object_2.Initialize();