Skip to content
Merged
Show file tree
Hide file tree
Changes from 9 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
8 changes: 8 additions & 0 deletions cmake/global_variables.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,14 @@ file(GLOB generator_srcs CONFIGURE_DEPENDS
"${GENERATORS_ROOT}/openvino/*.cpp"
"${GENERATORS_ROOT}/ryzenai/*.h"
"${GENERATORS_ROOT}/ryzenai/*.cpp"
"${GENERATORS_ROOT}/cuda/session_options.h"
"${GENERATORS_ROOT}/cuda/session_options.cpp"
"${GENERATORS_ROOT}/nvtrtrtx/*.h"
"${GENERATORS_ROOT}/nvtrtrtx/*.cpp"
"${GENERATORS_ROOT}/vitisai/*.h"
"${GENERATORS_ROOT}/vitisai/*.cpp"
"${GENERATORS_ROOT}/dml/session_options.h"
"${GENERATORS_ROOT}/dml/session_options.cpp"
"${MODELS_ROOT}/*.h"
"${MODELS_ROOT}/*.cpp"
"${ENGINE_ROOT}/*.h"
Expand Down
95 changes: 95 additions & 0 deletions src/cuda/session_options.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.

#include "session_options.h"
#include "../models/session_options.h"

namespace Generators::CUDAExecutionProvider {

namespace {

void AppendProviderBridgeExecutionProvider(
OrtSessionOptions& session_options,
const Config::ProviderOptions& provider_options,
DeviceInterface*& device) {
auto ort_provider_options = OrtCUDAProviderOptionsV2::Create();
std::vector<const char*> keys, values;

// Memory management settings
const char* arena_keys[] = {
"max_mem",
"arena_extend_strategy",
"initial_chunk_size_bytes",
"max_dead_bytes_per_chunk",
"initial_growth_chunk_size_bytes"};
size_t arena_values[] = {
static_cast<size_t>(0),
static_cast<size_t>(-1),
static_cast<size_t>(-1),
static_cast<size_t>(-1),
static_cast<size_t>(-1)};
bool use_arena_management = false;

for (auto& option : provider_options.options) {
auto it = std::find(std::begin(arena_keys), std::end(arena_keys), option.first);

if (it == std::end(arena_keys)) {
keys.emplace_back(option.first.c_str());
values.emplace_back(option.second.c_str());
} else {
const size_t idx = std::distance(std::begin(arena_keys), it);
long long parsed_value = std::stoll(option.second);
if (parsed_value < -1) {
throw std::out_of_range("Arena configuration option value is out of range");
}
arena_values[idx] = (parsed_value == -1)
? static_cast<size_t>(-1)
: static_cast<size_t>(parsed_value);
use_arena_management = true;
}
}
ort_provider_options->Update(keys.data(), values.data(), keys.size());

// Device type determines the scoring device.
// Create and set our cudaStream_t
ort_provider_options->UpdateValue("user_compute_stream", device->GetCudaStream());

// Use fine-grained memory management of BFC Arena
if (use_arena_management) {
auto arena_cfg = OrtArenaCfg::Create(arena_keys, arena_values, std::size(arena_keys));
ort_provider_options->UpdateValue("default_memory_arena_cfg", arena_cfg.get());
}

session_options.AppendExecutionProvider_CUDA_V2(*ort_provider_options);
Comment thread
baijumeswani marked this conversation as resolved.
}

} // namespace

void AddCudaStreamConfig(OrtSessionOptions& session_options, DeviceInterface* device,
const std::string& config_key) {
if (device) {
void* stream_ptr = device->GetCudaStream();
std::stringstream stream_value;
stream_value << reinterpret_cast<uintptr_t>(stream_ptr);
session_options.AddConfigEntry(config_key.c_str(), stream_value.str().c_str());
}
}

DeviceInterface* AppendExecutionProvider(OrtSessionOptions& session_options,
const Config::ProviderOptions& provider_options,
const Config& /*config*/,
bool /*disable_graph_capture*/) {
auto device = GetDeviceInterface(DeviceType::CUDA);
AddCudaStreamConfig(session_options, device);
// Try pre-registered plugin path first
if (!AppendExecutionProviderV2(session_options, provider_options,
DeviceType::CUDA, "CUDAExecutionProvider")) {
// Register the CUDA execution provider as a provider-bridge provider.
CUDAExecutionProvider::AppendProviderBridgeExecutionProvider(
session_options, provider_options, device);
}

return device;
}

} // namespace Generators::CUDAExecutionProvider
22 changes: 22 additions & 0 deletions src/cuda/session_options.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#pragma once

#include "../generators.h"

namespace Generators::CUDAExecutionProvider {

// Writes the CUDA compute stream pointer (as a stringified integer) into a
// session config entry keyed by |config_key|. This is used by both the CUDA
// and NvTensorRtRtx providers so that the EP can share the same stream.
void AddCudaStreamConfig(OrtSessionOptions& session_options, DeviceInterface* device,
const std::string& config_key = "user_compute_stream");

// Registers the CUDA execution provider on |session_options|. Tries the V2
// plugin path first; falls back to the provider-bridge (CUDA V2 options) path.
DeviceInterface* AppendExecutionProvider(OrtSessionOptions& session_options,
const Config::ProviderOptions& provider_options,
const Config& config,
bool disable_graph_capture = false);

} // namespace Generators::CUDAExecutionProvider
51 changes: 51 additions & 0 deletions src/dml/session_options.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.

#include "session_options.h"
#include "../models/session_options.h"

Comment thread
baijumeswani marked this conversation as resolved.
namespace Generators::DMLExecutionProvider {

DeviceInterface* AppendExecutionProvider(OrtSessionOptions& session_options,
const Config::ProviderOptions& provider_options,
const Config& /*config*/,
bool disable_graph_capture) {
#if USE_DML
if (!GetDmlInterface()) {
LUID device_luid{};
LUID* p_device_luid{};
uint32_t device_index{};
uint32_t* p_device_index{};
for (const auto& [name, value] : provider_options.options) {
if (name == "luid") {
if (auto separator_position = value.find(":"); separator_position != std::string::npos) {
device_luid.HighPart = std::stol(value.substr(0, separator_position));
device_luid.LowPart = std::stol(value.substr(separator_position + 1));
p_device_luid = &device_luid;
}
} else if (name == "device_index") {
device_index = std::stoi(value);
p_device_index = &device_index;
}
}

InitDmlInterface(p_device_luid, p_device_index);
}

// Non-decoder sessions (vision, speech, embedding) have control-flow nodes
// that are incompatible with graph capture, so the caller sets
// disable_graph_capture=true for those sessions.
if (!disable_graph_capture) {
session_options.AddConfigEntry("ep.dml.enable_graph_capture", "1");
}

SetDmlProvider(session_options);

auto device = GetDeviceInterface(DeviceType::DML); // We use a DML allocator for input/output caches, but other tensors will use CPU tensors
return device;
#else
throw std::runtime_error("DML provider requested, but the installed GenAI has not been built with DML support");
#endif
}

} // namespace Generators::DMLExecutionProvider
16 changes: 16 additions & 0 deletions src/dml/session_options.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#pragma once

#include "../generators.h"

namespace Generators::DMLExecutionProvider {

// Initialises the DML interface (if not already done), optionally enables graph
// capture, and registers the DirectML execution provider on |session_options|.
DeviceInterface* AppendExecutionProvider(OrtSessionOptions& session_options,
const Config::ProviderOptions& provider_options,
const Config& config,
bool disable_graph_capture = false);

} // namespace Generators::DMLExecutionProvider
2 changes: 1 addition & 1 deletion src/models/marian.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ namespace Generators {
MarianModel::MarianModel(std::unique_ptr<Config> config, OrtEnv& ort_env)
: Model{std::move(config)} {
encoder_session_options_ = OrtSessionOptions::Create();
CreateSessionOptionsFromConfig(config_->model.encoder.session_options.has_value() ? config_->model.encoder.session_options.value() : config_->model.decoder.session_options, *encoder_session_options_, true, false);
CreateSessionOptionsFromConfig(config_->model.encoder.session_options.has_value() ? config_->model.encoder.session_options.value() : config_->model.decoder.session_options, *encoder_session_options_, true);

session_encoder_ = CreateSession(ort_env, config_->model.encoder.filename, encoder_session_options_.get());
session_decoder_ = CreateSession(ort_env, config_->model.decoder.filename, session_options_.get());
Expand Down
Loading
Loading