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
16 changes: 12 additions & 4 deletions src/config.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,8 @@ struct ProviderOptionsObject_Element : JSON::Element {
};

struct ProviderOptionsArray_Element : JSON::Element {
explicit ProviderOptionsArray_Element(std::vector<Config::ProviderOptions>& v) : v_{v} {}
explicit ProviderOptionsArray_Element(std::vector<Config::ProviderOptions>& v, std::vector<std::string>& providers)
: v_{v}, providers_{providers} {}

JSON::Element& OnObject(std::string_view name) override { return object_; }

Expand All @@ -68,11 +69,18 @@ struct ProviderOptionsArray_Element : JSON::Element {
} else if (v.name == "dml") {
v.name = "DML";
}

if (std::find(providers_.begin(), providers_.end(), v.name) == providers_.end()) {
Comment thread
baijumeswani marked this conversation as resolved.
// 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<Config::ProviderOptions>& v_;
std::vector<std::string>& providers_;
ProviderOptionsObject_Element object_{v_};
};

Expand Down Expand Up @@ -143,7 +151,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_.providers};
NamedStrings_Element config_entries_{v_.config_entries};
};

Expand Down Expand Up @@ -711,7 +719,7 @@ 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();
Comment thread
baijumeswani marked this conversation as resolved.
}

void SetProviderOption(Config& config, std::string_view provider_name, std::string_view option_name, std::string_view option_value) {
Expand All @@ -721,7 +729,7 @@ void SetProviderOption(Config& config, std::string_view provider_name, std::stri
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());
}

Expand Down
1 change: 1 addition & 0 deletions src/config.h
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,7 @@ struct Config {
std::vector<NamedString> config_entries; // Entries go into OrtSessionOptions::AddConfigEntry

std::vector<ProviderOptions> provider_options;
std::vector<std::string> providers; // List of providers to use at runtime, not persisted in the json currently
std::optional<GraphOptimizationLevel> graph_optimization_level;
};

Expand Down
19 changes: 15 additions & 4 deletions src/models/model.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -274,12 +274,21 @@ int32_t Tokenizer::TokenToTokenId(const char* token) const {
}

DeviceInterface* SetProviderSessionOptions(OrtSessionOptions& session_options,
const std::vector<std::string>& providers,
const std::vector<Config::ProviderOptions>& 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()) {
Comment thread
baijumeswani marked this conversation as resolved.
throw std::runtime_error("Provider options not found for provider: " + provider);
}
const auto& provider_options = *provider_options_it;

if (provider_options.name == "cuda") {
auto ort_provider_options = OrtCUDAProviderOptionsV2::Create();
std::vector<const char*> keys, values;
Expand All @@ -299,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;

Expand Down Expand Up @@ -416,7 +424,8 @@ void EnsureDeviceOrtInit(DeviceInterface& device) {
auto session_options = OrtSessionOptions::Create();
std::vector<Config::ProviderOptions> provider_options_list;
provider_options_list.emplace_back(Config::ProviderOptions{device_type_names[static_cast<int>(type)], {}});
SetProviderSessionOptions(*session_options, provider_options_list, true, false);
const std::vector<std::string> providers{device_type_names[static_cast<int>(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());
Expand Down Expand Up @@ -613,7 +622,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_)
Expand Down