From 96f09e3c5fd218d3074fadec89d0dbf39a3bd336 Mon Sep 17 00:00:00 2001 From: Baiju Meswani Date: Tue, 6 May 2025 18:22:29 +0000 Subject: [PATCH 1/3] Persist provider options across clearproviders, appendprovider where possible --- src/config.cpp | 28 ++++++++++++++++++++++++++-- src/config.h | 1 + 2 files changed, 27 insertions(+), 2 deletions(-) diff --git a/src/config.cpp b/src/config.cpp index 4e346d1ac7..bf156db40a 100644 --- a/src/config.cpp +++ b/src/config.cpp @@ -54,7 +54,9 @@ struct ProviderOptionsObject_Element : JSON::Element { }; struct ProviderOptionsArray_Element : JSON::Element { - explicit ProviderOptionsArray_Element(std::vector& v) : v_{v} {} + explicit ProviderOptionsArray_Element(std::vector& v, + std::vector* v_backup = nullptr) + : v_{v}, v_backup_{v_backup} {} JSON::Element& OnObject(std::string_view name) override { return object_; } @@ -69,10 +71,18 @@ struct ProviderOptionsArray_Element : JSON::Element { v.name = "DML"; } } + + // Copy the provider options to the backup vector. + // The backup is used to restore the original provider options when + // AppendProvider is called with the same provider name. + if (v_backup_) { + *v_backup_ = v_; + } } private: std::vector& v_; + std::vector* v_backup_; ProviderOptionsObject_Element object_{v_}; }; @@ -143,7 +153,7 @@ struct SessionOptions_Element : JSON::Element { private: Config::SessionOptions& v_; - ProviderOptionsArray_Element provider_options_{v_.provider_options}; + ProviderOptionsArray_Element provider_options_{v_.provider_options, &v_.provider_options_backup}; NamedStrings_Element config_entries_{v_.config_entries}; }; @@ -715,6 +725,20 @@ void ClearProviders(Config& config) { } void SetProviderOption(Config& config, std::string_view provider_name, std::string_view option_name, std::string_view option_value) { + for (auto& provider_options : config.model.decoder.session_options.provider_options_backup) { + // If config.model.decoder.session_options.provider_options already has the provider name, add the requested option to it. + // Otherwise, if the provider name is found in the backup, query the provider options from the backup and add them to the current provider options. + // Otherwise, create a new provider options object and add it to the current provider options. + // This ensures that any __required__ options specified in the genai_config.json are not lost when ClearProviders -> AppendProvider is called. + // This is important for providers like Cuda, which require certain options to be set for some models. + if (provider_options.name == provider_name && + std::none_of(config.model.decoder.session_options.provider_options.begin(), + config.model.decoder.session_options.provider_options.end(), + [&provider_name](const Config::ProviderOptions& po) { return po.name == provider_name; })) { + config.model.decoder.session_options.provider_options.emplace_back(provider_options); + return; + } + } std::ostringstream json; json << R"({")" << provider_name << R"(":{)"; if (!option_name.empty()) { diff --git a/src/config.h b/src/config.h index ea18bd11c8..63a2c059cf 100644 --- a/src/config.h +++ b/src/config.h @@ -74,6 +74,7 @@ struct Config { std::vector config_entries; // Entries go into OrtSessionOptions::AddConfigEntry std::vector provider_options; + std::vector provider_options_backup; std::optional graph_optimization_level; }; From 63ec160ce9ee680edabb52c5d7d6f9a5fa73cca1 Mon Sep 17 00:00:00 2001 From: Baiju Meswani Date: Tue, 6 May 2025 19:45:12 +0000 Subject: [PATCH 2/3] Address pull request review comments --- src/config.cpp | 38 +++++++++++--------------------------- src/config.h | 2 +- src/models/model.cpp | 18 +++++++++++++++--- 3 files changed, 27 insertions(+), 31 deletions(-) diff --git a/src/config.cpp b/src/config.cpp index bf156db40a..4770d9566a 100644 --- a/src/config.cpp +++ b/src/config.cpp @@ -54,9 +54,8 @@ struct ProviderOptionsObject_Element : JSON::Element { }; struct ProviderOptionsArray_Element : JSON::Element { - explicit ProviderOptionsArray_Element(std::vector& v, - std::vector* v_backup = nullptr) - : v_{v}, v_backup_{v_backup} {} + explicit ProviderOptionsArray_Element(std::vector& v, std::vector& providers) + : v_{v}, providers_{providers} {} JSON::Element& OnObject(std::string_view name) override { return object_; } @@ -70,19 +69,18 @@ struct ProviderOptionsArray_Element : JSON::Element { } else if (v.name == "dml") { v.name = "DML"; } - } - // Copy the provider options to the backup vector. - // The backup is used to restore the original provider options when - // AppendProvider is called with the same provider name. - if (v_backup_) { - *v_backup_ = v_; + if (std::find(providers_.begin(), providers_.end(), v.name) == providers_.end()) { + // The providers array determines the which execution provider is picked for the session.. + // It also determines the order of the providers. + providers_.push_back(v.name); + } } } private: std::vector& v_; - std::vector* v_backup_; + std::vector& providers_; ProviderOptionsObject_Element object_{v_}; }; @@ -153,7 +151,7 @@ struct SessionOptions_Element : JSON::Element { private: Config::SessionOptions& v_; - ProviderOptionsArray_Element provider_options_{v_.provider_options, &v_.provider_options_backup}; + ProviderOptionsArray_Element provider_options_{v_.provider_options, v_.providers}; NamedStrings_Element config_entries_{v_.config_entries}; }; @@ -721,31 +719,17 @@ void SetSearchBool(Config::Search& search, std::string_view name, bool value) { } void ClearProviders(Config& config) { - config.model.decoder.session_options.provider_options.clear(); + config.model.decoder.session_options.providers.clear(); } void SetProviderOption(Config& config, std::string_view provider_name, std::string_view option_name, std::string_view option_value) { - for (auto& provider_options : config.model.decoder.session_options.provider_options_backup) { - // If config.model.decoder.session_options.provider_options already has the provider name, add the requested option to it. - // Otherwise, if the provider name is found in the backup, query the provider options from the backup and add them to the current provider options. - // Otherwise, create a new provider options object and add it to the current provider options. - // This ensures that any __required__ options specified in the genai_config.json are not lost when ClearProviders -> AppendProvider is called. - // This is important for providers like Cuda, which require certain options to be set for some models. - if (provider_options.name == provider_name && - std::none_of(config.model.decoder.session_options.provider_options.begin(), - config.model.decoder.session_options.provider_options.end(), - [&provider_name](const Config::ProviderOptions& po) { return po.name == provider_name; })) { - config.model.decoder.session_options.provider_options.emplace_back(provider_options); - return; - } - } std::ostringstream json; json << R"({")" << provider_name << R"(":{)"; if (!option_name.empty()) { json << R"(")" << option_name << R"(":")" << option_value << R"(")"; } json << R"(}})"; - ProviderOptionsArray_Element element{config.model.decoder.session_options.provider_options}; + ProviderOptionsArray_Element element{config.model.decoder.session_options.provider_options, config.model.decoder.session_options.providers}; JSON::Parse(element, json.str()); } diff --git a/src/config.h b/src/config.h index 63a2c059cf..882afb2e0b 100644 --- a/src/config.h +++ b/src/config.h @@ -74,7 +74,7 @@ struct Config { std::vector config_entries; // Entries go into OrtSessionOptions::AddConfigEntry std::vector provider_options; - std::vector provider_options_backup; + std::vector providers; std::optional graph_optimization_level; }; diff --git a/src/models/model.cpp b/src/models/model.cpp index 4ccf866c8f..91fbec259b 100644 --- a/src/models/model.cpp +++ b/src/models/model.cpp @@ -274,12 +274,21 @@ int32_t Tokenizer::TokenToTokenId(const char* token) const { } DeviceInterface* SetProviderSessionOptions(OrtSessionOptions& session_options, + const std::vector& providers, const std::vector& provider_options_list, bool is_primary_session_options, bool disable_graph_capture) { DeviceInterface* p_device{}; - for (auto& provider_options : provider_options_list) { + for (auto& provider : providers) { + auto provider_options_it = std::find_if(provider_options_list.begin(), provider_options_list.end(), + [&provider](const Config::ProviderOptions& po) { return po.name == provider; }); + + if (provider_options_it == provider_options_list.end()) { + throw std::runtime_error("Provider options not found for provider: " + provider); + } + auto provider_options = *provider_options_it; + if (provider_options.name == "cuda") { auto ort_provider_options = OrtCUDAProviderOptionsV2::Create(); std::vector keys, values; @@ -416,7 +425,8 @@ void EnsureDeviceOrtInit(DeviceInterface& device) { auto session_options = OrtSessionOptions::Create(); std::vector provider_options_list; provider_options_list.emplace_back(Config::ProviderOptions{device_type_names[static_cast(type)], {}}); - SetProviderSessionOptions(*session_options, provider_options_list, true, false); + const std::vector providers{device_type_names[static_cast(type)]}; + SetProviderSessionOptions(*session_options, providers, provider_options_list, true, false); session_options->SetLogSeverityLevel(ORT_LOGGING_LEVEL_ERROR); // Errors only here, as warnings are not useful to the user allocator.session_ = OrtSession::Create(GetOrtEnv(), g_trivial_model, sizeof(g_trivial_model), session_options.get()); @@ -613,7 +623,9 @@ void Model::CreateSessionOptionsFromConfig(const Config::SessionOptions& config_ session_options.SetGraphOptimizationLevel(config_session_options.graph_optimization_level.value()); } - p_device_ = SetProviderSessionOptions(session_options, config_session_options.provider_options, is_primary_session_options, disable_graph_capture); + p_device_ = SetProviderSessionOptions(session_options, config_session_options.providers, + config_session_options.provider_options, is_primary_session_options, + disable_graph_capture); // Fallback to CPU if no provider specific interface was set if (!p_device_) From b050d20921ba3d04fabbdc1b747f6f73b7ff7bb8 Mon Sep 17 00:00:00 2001 From: Baiju Meswani Date: Wed, 7 May 2025 02:39:22 +0000 Subject: [PATCH 3/3] Address pr review comments --- src/config.h | 2 +- src/models/model.cpp | 3 +-- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/src/config.h b/src/config.h index 882afb2e0b..cf406b3709 100644 --- a/src/config.h +++ b/src/config.h @@ -74,7 +74,7 @@ struct Config { std::vector config_entries; // Entries go into OrtSessionOptions::AddConfigEntry std::vector provider_options; - std::vector providers; + std::vector providers; // List of providers to use at runtime, not persisted in the json currently std::optional graph_optimization_level; }; diff --git a/src/models/model.cpp b/src/models/model.cpp index 91fbec259b..fcd0f3d6fa 100644 --- a/src/models/model.cpp +++ b/src/models/model.cpp @@ -287,7 +287,7 @@ DeviceInterface* SetProviderSessionOptions(OrtSessionOptions& session_options, if (provider_options_it == provider_options_list.end()) { throw std::runtime_error("Provider options not found for provider: " + provider); } - auto provider_options = *provider_options_it; + const auto& provider_options = *provider_options_it; if (provider_options.name == "cuda") { auto ort_provider_options = OrtCUDAProviderOptionsV2::Create(); @@ -308,7 +308,6 @@ DeviceInterface* SetProviderSessionOptions(OrtSessionOptions& session_options, } session_options.AppendExecutionProvider_CUDA_V2(*ort_provider_options); - } else if (provider_options.name == "rocm") { OrtROCMProviderOptions ort_provider_options;