From 29dcfb24ab13ffa585131acd1f623f7ca5da073b Mon Sep 17 00:00:00 2001 From: Pranav Sharma Date: Sat, 22 Aug 2020 10:03:17 -0700 Subject: [PATCH 1/4] Allow multiple sessions to share an allocator, optimize constant folding memory usage, expose arena configs. (#4813) * Add support for sharing allocators * Incremental update * Address some PR comments, add unit tests, add documentation. * Address PR comments, add tests and some documentation. * Fix build and test issues * Remove RegisterAllocator API restoring the OrtAllocator interface changes. Changed docs to reflect this. Also fixed the orttraining segfault. The segfault was because in the case of training session, the CPU exec prov is not available at the time the transformers are applied. Changed it to create a new one. --- docs/C_API.md | 7 ++ .../core/framework/execution_provider.h | 11 +- .../core/optimizer/graph_transformer_utils.h | 2 + .../onnxruntime/core/session/environment.h | 15 +++ .../core/session/onnxruntime_c_api.h | 23 ++++ .../core/session/onnxruntime_cxx_api.h | 2 + .../core/session/onnxruntime_cxx_inline.h | 5 + .../onnxruntime_session_options_config_keys.h | 6 +- onnxruntime/core/framework/allocatormgr.cc | 6 +- onnxruntime/core/framework/allocatormgr.h | 13 +- onnxruntime/core/framework/bfc_arena.cc | 27 ++-- onnxruntime/core/framework/bfc_arena.h | 11 +- .../core/framework/execution_provider.cc | 19 ++- .../core/framework/execution_providers.h | 9 ++ .../core/optimizer/constant_folding.cc | 13 +- onnxruntime/core/optimizer/constant_folding.h | 15 ++- .../core/optimizer/graph_transformer_utils.cc | 4 +- .../optimizer/optimizer_execution_frame.cc | 12 +- .../optimizer/optimizer_execution_frame.h | 6 +- onnxruntime/core/session/allocator_impl.h | 3 + onnxruntime/core/session/device_allocator.cc | 117 ++++++++++++------ onnxruntime/core/session/device_allocator.h | 45 +++++++ onnxruntime/core/session/environment.cc | 14 +++ onnxruntime/core/session/inference_session.cc | 52 ++++++-- onnxruntime/core/session/inference_session.h | 11 +- onnxruntime/core/session/onnxruntime_c_api.cc | 2 +- onnxruntime/core/session/ort_apis.h | 1 + onnxruntime/core/session/ort_env.cc | 6 + onnxruntime/core/session/ort_env.h | 7 ++ .../test/framework/inference_session_test.cc | 106 ++++++++++++++++ onnxruntime/test/optimizer/cse_test.cc | 3 +- .../test/optimizer/graph_transform_test.cc | 23 ++-- .../optimizer/graph_transform_utils_test.cc | 13 +- onnxruntime/test/optimizer/optimizer_test.cc | 2 +- onnxruntime/test/shared_lib/test_inference.cc | 55 ++++++++ .../core/optimizer/graph_transformer_utils.cc | 3 +- .../core/optimizer/graph_transformer_utils.h | 1 + .../core/session/training_session.cc | 16 ++- .../core/session/training_session.h | 3 +- 39 files changed, 575 insertions(+), 114 deletions(-) create mode 100644 onnxruntime/core/session/device_allocator.h diff --git a/docs/C_API.md b/docs/C_API.md index b6e078d1ec..b7266015c3 100644 --- a/docs/C_API.md +++ b/docs/C_API.md @@ -21,6 +21,13 @@ is as follows * Create env using ```CreateEnvWithGlobalThreadPools()``` * Create session and call ```DisablePerSessionThreads()``` on the session options object * Call ```Run()``` as usual +* **Share allocator(s) between sessions:** Allow multiple sessions in the same process to use the same allocator(s). This +allocator is first registered in the env and then reused by all sessions that use the same env instance unless a session +chooses to override this by setting ```session_state.use_env_allocators``` to "0". Usage of this feature is as follows + * Register an allocator created by ORT using the ```CreateAndRegisterAllocator``` API. + * Set ```session.use_env_allocators``` to "1" for each session that wants to use the env registered allocators. + * See test ```TestSharedAllocatorUsingCreateAndRegisterAllocator``` in + onnxruntime/test/shared_lib/test_inference.cc for an example. ## Usage Overview diff --git a/include/onnxruntime/core/framework/execution_provider.h b/include/onnxruntime/core/framework/execution_provider.h index 15b0683ece..654487bfde 100644 --- a/include/onnxruntime/core/framework/execution_provider.h +++ b/include/onnxruntime/core/framework/execution_provider.h @@ -22,7 +22,8 @@ class KernelRegistryManager; /** Logical device representation. */ -typedef std::map AllocatorMap; +using AllocatorMap = std::map; +using MemoryInfoSet = std::set; // if we are export the fused function to dll, the function will still in the same binary as onnxruntime // use std function to give execution provider some chance to capture some state. @@ -34,8 +35,8 @@ using DestroyFunctionStateFunc = std::function; using UnorderedMapStringToString = std::unordered_map; //data types for execution provider options -using ProviderOptionsVector = std::vector; -using ProviderOptionsMap = std::unordered_map; +using ProviderOptionsVector = std::vector; +using ProviderOptionsMap = std::unordered_map; struct NodeComputeInfo { CreateFunctionStateFunc create_state_func; @@ -113,7 +114,7 @@ class IExecutionProvider { /** Store execution provider's configurations. */ - void SetProviderOptions(UnorderedMapStringToString& options) { + void SetProviderOptions(UnorderedMapStringToString& options) { provider_options_ = options; } @@ -165,6 +166,7 @@ class IExecutionProvider { virtual common::Status OnSessionInitializationEnd(); void InsertAllocator(AllocatorPtr allocator); + void ReplaceAllocator(AllocatorPtr allocator); /** Given a list of fused_node, return create_state/compute/release_state func for each node. @@ -193,6 +195,7 @@ class IExecutionProvider { private: const std::string type_; AllocatorMap allocators_; + MemoryInfoSet mem_info_set_; // to ensure only allocators with unique OrtMemoryInfo are registered in the provider. //It will be set when this object is registered to a session const logging::Logger* logger_ = nullptr; // convenience list of the allocators so GetAllocatorList doesn't have to build a new vector each time diff --git a/include/onnxruntime/core/optimizer/graph_transformer_utils.h b/include/onnxruntime/core/optimizer/graph_transformer_utils.h index 1d2b430f34..b5d0287f5d 100644 --- a/include/onnxruntime/core/optimizer/graph_transformer_utils.h +++ b/include/onnxruntime/core/optimizer/graph_transformer_utils.h @@ -11,6 +11,7 @@ namespace onnxruntime { struct FreeDimensionOverride; +class IExecutionProvider; namespace optimizer_utils { @@ -25,6 +26,7 @@ std::vector> GenerateRewriteRules(TransformerLevel and the transformers_and_rules_to_enable. */ std::vector> GenerateTransformers(TransformerLevel level, gsl::span free_dimension_overrides, + const IExecutionProvider& execution_provider /*required by constant folding*/, const std::vector& rules_and_transformers_to_enable = {}); /** Given a TransformerLevel, this method generates a name for the rule-based graph transformer of that level. */ diff --git a/include/onnxruntime/core/session/environment.h b/include/onnxruntime/core/session/environment.h index 795cdb3672..1a2ef85caa 100644 --- a/include/onnxruntime/core/session/environment.h +++ b/include/onnxruntime/core/session/environment.h @@ -9,6 +9,7 @@ #include "core/common/status.h" #include "core/platform/threadpool.h" #include "core/common/logging/logging.h" +#include "core/framework/allocator.h" struct OrtThreadingOptions; namespace onnxruntime { @@ -54,6 +55,19 @@ class Environment { return create_global_thread_pools_; } + /** + * Registers an allocator for sharing between multiple sessions. + * Return an error if an allocator with the same OrtMemoryInfo is already registered. + */ + Status RegisterAllocator(AllocatorPtr allocator); + + /** + * Returns the list of registered allocators in this env. + */ + const std::vector& GetRegisteredSharedAllocators() const { + return shared_allocators_; + } + private: ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(Environment); @@ -66,5 +80,6 @@ class Environment { std::unique_ptr intra_op_thread_pool_; std::unique_ptr inter_op_thread_pool_; bool create_global_thread_pools_{false}; + std::vector shared_allocators_; }; } // namespace onnxruntime diff --git a/include/onnxruntime/core/session/onnxruntime_c_api.h b/include/onnxruntime/core/session/onnxruntime_c_api.h index 5e3fa3c58c..852ce7dd1f 100644 --- a/include/onnxruntime/core/session/onnxruntime_c_api.h +++ b/include/onnxruntime/core/session/onnxruntime_c_api.h @@ -144,6 +144,16 @@ typedef enum OrtErrorCode { ORT_EP_FAIL, } OrtErrorCode; +// This configures the arena based allocator used by ORT +// See ONNX_Runtime_Perf_Tuning.md for details on what these mean and how to choose these values +// Use -1 to allow ORT to choose defaults for all the options below +typedef struct OrtArenaCfg { + int max_mem; + int arena_extend_strategy; // 0 = kNextPowerOfTwo, 1 = kSameAsRequested + int initial_chunk_size_bytes; + int max_dead_bytes_per_chunk; +} OrtArenaCfg; + #define ORT_RUNTIME_CLASS(X) \ struct Ort##X; \ typedef struct Ort##X Ort##X; @@ -578,6 +588,7 @@ struct OrtApi { // The returned pointer doesn't have to be freed. // Always returns the same instance on every invocation. + // Please note that this is a non-arena based allocator. ORT_API2_STATUS(GetAllocatorWithDefaultOptions, _Outptr_ OrtAllocator** out); // Override symbolic dimensions (by specific denotation strings) with actual values if known at session initialization time to enable @@ -990,6 +1001,18 @@ struct OrtApi { * This is a no-copy method whose pointer is only valid until the backing OrtValue is free'd. */ ORT_API2_STATUS(TensorAt, _Inout_ OrtValue* value, size_t* location_values, size_t location_values_count, _Outptr_ void** out); + + /** + * Creates an allocator instance and registers it with the env to enable + * sharing between multiple sessions that use the same env instance. + * Lifetime of the created allocator will be valid for the duration of the environment. + * Returns an error if an allocator with the same OrtMemoryInfo is already registered. + * \param mem_info must be non-null. + * \param arena_cfg if nullptr defaults will be used. + * See docs/C_API.md for details. + */ + ORT_API2_STATUS(CreateAndRegisterAllocator, _Inout_ OrtEnv* env, _In_ const OrtMemoryInfo* mem_info, + _In_ const OrtArenaCfg* arena_cfg); }; /* diff --git a/include/onnxruntime/core/session/onnxruntime_cxx_api.h b/include/onnxruntime/core/session/onnxruntime_cxx_api.h index 9b0ada9e46..9f2898618d 100644 --- a/include/onnxruntime/core/session/onnxruntime_cxx_api.h +++ b/include/onnxruntime/core/session/onnxruntime_cxx_api.h @@ -164,6 +164,8 @@ struct Env : Base { Env& EnableTelemetryEvents(); Env& DisableTelemetryEvents(); + Env& CreateAndRegisterAllocator(const OrtMemoryInfo* mem_info, const OrtArenaCfg* arena_cfg); + static const OrtApi* s_api; }; diff --git a/include/onnxruntime/core/session/onnxruntime_cxx_inline.h b/include/onnxruntime/core/session/onnxruntime_cxx_inline.h index 1cffc2cd14..c41b569bc1 100644 --- a/include/onnxruntime/core/session/onnxruntime_cxx_inline.h +++ b/include/onnxruntime/core/session/onnxruntime_cxx_inline.h @@ -308,6 +308,11 @@ inline Env& Env::DisableTelemetryEvents() { return *this; } +inline Env& Env::CreateAndRegisterAllocator(const OrtMemoryInfo* mem_info, const OrtArenaCfg* arena_cfg) { + ThrowOnError(GetApi().CreateAndRegisterAllocator(p_, mem_info, arena_cfg)); + return *this; +} + inline CustomOpDomain::CustomOpDomain(const char* domain) { ThrowOnError(GetApi().CreateCustomOpDomain(domain, &p_)); } diff --git a/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h b/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h index e3de40d5b4..cec659500d 100644 --- a/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h +++ b/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h @@ -18,4 +18,8 @@ // Key for disable PrePacking, // If the config value is set to "1" then the prepacking is disabled, otherwise prepacking is enabled (default value) -#define ORT_SESSION_OPTIONS_CONFIG_DISABLEPREPACKING "session_state.disable_prepacking" +#define ORT_SESSION_OPTIONS_CONFIG_DISABLEPREPACKING "session.disable_prepacking" + +// A value of "1" means allocators registered in the env will be used. "0" means the allocators created in the session +// will be used. Use this to override the usage of env allocators on a per session level. +#define ORT_SESSION_OPTIONS_CONFIG_USE_ENV_ALLOCATORS "session.use_env_allocators" diff --git a/onnxruntime/core/framework/allocatormgr.cc b/onnxruntime/core/framework/allocatormgr.cc index 147452432c..8b6bb47e3d 100644 --- a/onnxruntime/core/framework/allocatormgr.cc +++ b/onnxruntime/core/framework/allocatormgr.cc @@ -22,7 +22,11 @@ AllocatorPtr CreateAllocator(const DeviceAllocatorRegistrationInfo& info, onnxruntime::make_unique(std::move(device_allocator), info.max_mem)); #else return std::shared_ptr( - onnxruntime::make_unique(std::move(device_allocator), info.max_mem, info.arena_extend_strategy)); + onnxruntime::make_unique(std::move(device_allocator), + info.max_mem, + info.arena_extend_strategy, + info.initial_chunk_size_bytes, + info.max_dead_bytes_per_chunk)); #endif } diff --git a/onnxruntime/core/framework/allocatormgr.h b/onnxruntime/core/framework/allocatormgr.h index 0126625746..c23c954da8 100644 --- a/onnxruntime/core/framework/allocatormgr.h +++ b/onnxruntime/core/framework/allocatormgr.h @@ -11,21 +11,30 @@ namespace onnxruntime { using DeviceAllocatorFactory = std::function(OrtDevice::DeviceId)>; +// TODO why does DeviceAllocatorRegistrationInfo have arena related configs? +// TODO even if it should, they should be inside their own struct (OrtArenaCfg) as opposed to +// littering them as individual members of DeviceAllocatorRegistrationInfo struct DeviceAllocatorRegistrationInfo { DeviceAllocatorRegistrationInfo(OrtMemType ort_mem_type, DeviceAllocatorFactory alloc_factory, size_t mem, - ArenaExtendStrategy strategy = ArenaExtendStrategy::kNextPowerOfTwo) + ArenaExtendStrategy strategy = BFCArena::DEFAULT_ARENA_EXTEND_STRATEGY, + int initial_chunk_size_bytes0 = BFCArena::DEFAULT_INITIAL_CHUNK_SIZE_BYTES, + int max_dead_bytes_per_chunk0 = BFCArena::DEFAULT_MAX_DEAD_BYTES_PER_CHUNK) : mem_type(ort_mem_type), factory(alloc_factory), max_mem(mem), - arena_extend_strategy(strategy) { + arena_extend_strategy(strategy), + initial_chunk_size_bytes(initial_chunk_size_bytes0), + max_dead_bytes_per_chunk(max_dead_bytes_per_chunk0) { } OrtMemType mem_type; DeviceAllocatorFactory factory; size_t max_mem; ArenaExtendStrategy arena_extend_strategy; + int initial_chunk_size_bytes; + int max_dead_bytes_per_chunk; }; AllocatorPtr CreateAllocator(const DeviceAllocatorRegistrationInfo& info, OrtDevice::DeviceId device_id = 0, diff --git a/onnxruntime/core/framework/bfc_arena.cc b/onnxruntime/core/framework/bfc_arena.cc index 80f7d71a93..2ad7189c93 100644 --- a/onnxruntime/core/framework/bfc_arena.cc +++ b/onnxruntime/core/framework/bfc_arena.cc @@ -2,11 +2,14 @@ // Licensed under the MIT License. #include "core/framework/bfc_arena.h" +#include namespace onnxruntime { BFCArena::BFCArena(std::unique_ptr resource_allocator, size_t total_memory, - ArenaExtendStrategy arena_extend_strategy) + ArenaExtendStrategy arena_extend_strategy, + int initial_chunk_size_bytes, + int max_dead_bytes_per_chunk) : IArenaAllocator(OrtMemoryInfo(resource_allocator->Info().name, OrtAllocatorType::OrtArenaAllocator, resource_allocator->Info().device, @@ -14,14 +17,17 @@ BFCArena::BFCArena(std::unique_ptr resource_allocator, resource_allocator->Info().mem_type)), device_allocator_(std::move(resource_allocator)), free_chunks_list_(kInvalidChunkHandle), - next_allocation_id_(1) { - LOGS_DEFAULT(INFO) << "Creating BFCArena for " << device_allocator_->Info().name; - - // TODO - consider to make the initial chunk size and max 'fragmentation' (kMaxDeadBytesInChunk) values configurable. - // But first we need to add a mechanism to allow that sort of low level configuration to be done - // without adding separate parameters to SessionOptions for every single one of them. - curr_region_allocation_bytes_ = RoundedBytes(std::min(total_memory, size_t{1048576})); + next_allocation_id_(1), + initial_chunk_size_bytes_(initial_chunk_size_bytes), + max_dead_bytes_per_chunk_(max_dead_bytes_per_chunk) { + LOGS_DEFAULT(INFO) << "Creating BFCArena for " << device_allocator_->Info().name + << " with following configs: initial_chunk_size_bytes: " << initial_chunk_size_bytes_ + << " max_dead_bytes_per_chunk: " << max_dead_bytes_per_chunk_ + << " memory limit: " << total_memory + << " arena_extend_strategy " << static_cast(arena_extend_strategy); + // static_cast>(arena_extend_strategy); doesn't work on this compiler + curr_region_allocation_bytes_ = RoundedBytes(std::min(total_memory, static_cast(initial_chunk_size_bytes_))); // Allocate the requested amount of memory. memory_limit_ = total_memory; stats_.bytes_limit = static_cast(total_memory); @@ -312,10 +318,9 @@ void* BFCArena::FindChunkPtr(BinNum bin_num, size_t rounded_bytes, // If we can break the size of the chunk into two reasonably large // pieces, do so. In any case don't waste more than - // kMaxDeadBytesInChunk bytes on padding this alloc. - const int64_t kMaxDeadBytesInChunk = 128 << 20; // 128mb + // max_dead_bytes_per_chunk bytes on padding this alloc. if (chunk->size >= rounded_bytes * 2 || - static_cast(chunk->size) - rounded_bytes >= kMaxDeadBytesInChunk) { + static_cast(chunk->size) - static_cast(rounded_bytes) >= max_dead_bytes_per_chunk_) { SplitChunk(h, rounded_bytes); chunk = ChunkFromHandle(h); // Update chunk pointer in case it moved } diff --git a/onnxruntime/core/framework/bfc_arena.h b/onnxruntime/core/framework/bfc_arena.h index 0fdd61a4bd..5606559720 100644 --- a/onnxruntime/core/framework/bfc_arena.h +++ b/onnxruntime/core/framework/bfc_arena.h @@ -54,9 +54,15 @@ enum class ArenaExtendStrategy : int32_t { // all requests to allocate memory go through this interface. class BFCArena : public IArenaAllocator { public: + static const ArenaExtendStrategy DEFAULT_ARENA_EXTEND_STRATEGY = ArenaExtendStrategy::kNextPowerOfTwo; + static const int DEFAULT_INITIAL_CHUNK_SIZE_BYTES = 1048576; + static const int DEFAULT_MAX_DEAD_BYTES_PER_CHUNK = 128 * 1024 * 1024; + BFCArena(std::unique_ptr resource_allocator, size_t total_memory, - ArenaExtendStrategy arena_extend_strategy = ArenaExtendStrategy::kNextPowerOfTwo); + ArenaExtendStrategy arena_extend_strategy = DEFAULT_ARENA_EXTEND_STRATEGY, + int initial_chunk_size_bytes = DEFAULT_INITIAL_CHUNK_SIZE_BYTES, + int max_dead_bytes_per_chunk = DEFAULT_MAX_DEAD_BYTES_PER_CHUNK); ~BFCArena() override; @@ -443,6 +449,9 @@ class BFCArena : public IArenaAllocator { std::unordered_map reserved_chunks_; + const int initial_chunk_size_bytes_; + const int max_dead_bytes_per_chunk_; + ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(BFCArena); }; #ifdef __GNUC__ diff --git a/onnxruntime/core/framework/execution_provider.cc b/onnxruntime/core/framework/execution_provider.cc index 0949b9c42d..4c34291315 100644 --- a/onnxruntime/core/framework/execution_provider.cc +++ b/onnxruntime/core/framework/execution_provider.cc @@ -51,14 +51,25 @@ common::Status IExecutionProvider::OnRunEnd() { return Status::OK(); } common::Status IExecutionProvider::OnSessionInitializationEnd() { return Status::OK(); } +// Update allocator in the provider if already present; ignore if not. +void IExecutionProvider::ReplaceAllocator(AllocatorPtr allocator) { + const auto& info = allocator->Info(); + auto ite = mem_info_set_.find(info); + if (ite != mem_info_set_.end()) { + const int key = MakeKey(info.id, info.mem_type); + allocators_[key] = allocator; + } +} + void IExecutionProvider::InsertAllocator(AllocatorPtr allocator) { const OrtMemoryInfo& info = allocator->Info(); - const int key = MakeKey(info.id, info.mem_type); - auto iter = allocators_.find(key); - if (iter != allocators_.end()) { + auto ite = mem_info_set_.find(info); + if (ite != mem_info_set_.end()) { ORT_THROW("duplicated allocator"); } - allocators_.insert(iter, {key, allocator}); + const int key = MakeKey(info.id, info.mem_type); + allocators_.insert({key, allocator}); + mem_info_set_.insert(ite, info); allocator_list_.push_back(allocator); } diff --git a/onnxruntime/core/framework/execution_providers.h b/onnxruntime/core/framework/execution_providers.h index 5de21839b1..8125bdd874 100644 --- a/onnxruntime/core/framework/execution_providers.h +++ b/onnxruntime/core/framework/execution_providers.h @@ -56,6 +56,15 @@ class ExecutionProviders { return exec_providers_[it->second].get(); } + IExecutionProvider* Get(onnxruntime::ProviderType provider_id) { + auto it = provider_idx_map_.find(provider_id); + if (it == provider_idx_map_.end()) { + return nullptr; + } + + return exec_providers_[it->second].get(); + } + bool Empty() const { return exec_providers_.empty(); } size_t NumProviders() const { return exec_providers_.size(); } diff --git a/onnxruntime/core/optimizer/constant_folding.cc b/onnxruntime/core/optimizer/constant_folding.cc index dd463cea60..0a7062434b 100644 --- a/onnxruntime/core/optimizer/constant_folding.cc +++ b/onnxruntime/core/optimizer/constant_folding.cc @@ -12,6 +12,14 @@ using namespace onnxruntime::common; namespace onnxruntime { +ConstantFolding::ConstantFolding(const IExecutionProvider& execution_provider, + const std::unordered_set& compatible_execution_providers, + const std::unordered_set& excluded_initializers) noexcept + : GraphTransformer("ConstantFolding", compatible_execution_providers), + excluded_initializers_(excluded_initializers), + execution_provider_(execution_provider) { +} + // We need to handle a Shape node separately as the input doesn't need to be a constant initializer for // Shape to be able to be constant folded. static bool ConstantFoldShapeNode(Graph& graph, Node& node) { @@ -96,11 +104,8 @@ Status ConstantFolding::ApplyImpl(Graph& graph, bool& modified, int graph_level, } // Create execution frame for executing constant nodes. - std::unique_ptr cpu_execution_provider = - onnxruntime::make_unique(CPUExecutionProviderInfo()); - // Create execution frame for executing constant nodes. - OptimizerExecutionFrame::Info info({node}, constant_inputs, std::move(cpu_execution_provider)); + OptimizerExecutionFrame::Info info({node}, constant_inputs, execution_provider_); std::vector fetch_mlvalue_idxs; for (const auto* node_out : node->OutputDefs()) { diff --git a/onnxruntime/core/optimizer/constant_folding.h b/onnxruntime/core/optimizer/constant_folding.h index 7306cabdc4..e4b3a7a483 100644 --- a/onnxruntime/core/optimizer/constant_folding.h +++ b/onnxruntime/core/optimizer/constant_folding.h @@ -5,6 +5,8 @@ #include "core/optimizer/graph_transformer.h" #include "core/framework/ml_value.h" +#include +#include "core/framework/execution_provider.h" namespace onnxruntime { @@ -16,16 +18,19 @@ it statically computes parts of the graph that rely only on constant initializer */ class ConstantFolding : public GraphTransformer { public: - /** Constant folding will not be applied to nodes that have one of initializers from excluded_initializers as input. - For pre-training, the trainable weights are those initializers to be excluded. */ - ConstantFolding(const std::unordered_set& compatible_execution_providers = {}, - const std::unordered_set& excluded_initializers = {}) noexcept - : GraphTransformer("ConstantFolding", compatible_execution_providers), excluded_initializers_(excluded_initializers) {} + /*! Constant folding will not be applied to nodes that have one of initializers from excluded_initializers as input. + For pre-training, the trainable weights are those initializers to be excluded. + \param execution_provider Execution provider instance to execute constant folding. + */ + ConstantFolding(const IExecutionProvider& execution_provider, + const std::unordered_set& compatible_execution_providers = {}, + const std::unordered_set& excluded_initializers = {}) noexcept; private: Status ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const override; const std::unordered_set excluded_initializers_; + const IExecutionProvider& execution_provider_; }; } // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/graph_transformer_utils.cc b/onnxruntime/core/optimizer/graph_transformer_utils.cc index c6ab466338..a8a1e144cb 100644 --- a/onnxruntime/core/optimizer/graph_transformer_utils.cc +++ b/onnxruntime/core/optimizer/graph_transformer_utils.cc @@ -36,6 +36,7 @@ #include "core/optimizer/unsqueeze_elimination.h" namespace onnxruntime { +class IExecutionProvider; namespace optimizer_utils { @@ -107,6 +108,7 @@ std::unique_ptr GenerateRuleBasedGraphTransformer(Tra std::vector> GenerateTransformers(TransformerLevel level, gsl::span free_dimension_overrides, + const IExecutionProvider& execution_provider, /*required by constant folding*/ const std::vector& transformers_and_rules_to_enable) { std::vector> transformers; std::unique_ptr rule_transformer = nullptr; @@ -115,7 +117,7 @@ std::vector> GenerateTransformers(TransformerL std::unordered_set l1_execution_providers = {}; transformers.emplace_back(onnxruntime::make_unique(l1_execution_providers)); - transformers.emplace_back(onnxruntime::make_unique(l1_execution_providers)); + transformers.emplace_back(onnxruntime::make_unique(execution_provider, l1_execution_providers)); transformers.emplace_back(onnxruntime::make_unique(l1_execution_providers)); transformers.emplace_back(onnxruntime::make_unique(l1_execution_providers)); transformers.emplace_back(onnxruntime::make_unique(free_dimension_overrides)); diff --git a/onnxruntime/core/optimizer/optimizer_execution_frame.cc b/onnxruntime/core/optimizer/optimizer_execution_frame.cc index 8ff9d325ac..3c7a93da01 100644 --- a/onnxruntime/core/optimizer/optimizer_execution_frame.cc +++ b/onnxruntime/core/optimizer/optimizer_execution_frame.cc @@ -19,11 +19,9 @@ namespace onnxruntime { OptimizerExecutionFrame::Info::Info(const std::vector& nodes, const InitializedTensorSet& initialized_tensor_set, - std::unique_ptr cpu_execution_provider) { - ORT_ENFORCE(cpu_execution_provider, "Provided CPU execution provider is a nullptr"); - cpu_execution_provider_ = std::move(cpu_execution_provider); - - allocator_ptr_ = cpu_execution_provider_->GetAllocator(device_id_, mem_type_); + const IExecutionProvider& execution_provider) + : execution_provider_(execution_provider) { + allocator_ptr_ = execution_provider_.GetAllocator(device_id_, mem_type_); ORT_ENFORCE(allocator_ptr_, "Failed to get allocator for optimizer"); data_transfer_mgr_.RegisterDataTransfer(onnxruntime::make_unique()); @@ -67,8 +65,8 @@ OptimizerExecutionFrame::Info::Info(const std::vector& nodes, std::unique_ptr OptimizerExecutionFrame::Info::CreateKernel(const Node* node) const { std::unique_ptr op_kernel; - std::shared_ptr kernel_registry = cpu_execution_provider_->GetKernelRegistry(); - auto status = kernel_registry->TryCreateKernel(*node, *cpu_execution_provider_, initializers_, + std::shared_ptr kernel_registry = execution_provider_.GetKernelRegistry(); + auto status = kernel_registry->TryCreateKernel(*node, execution_provider_, initializers_, ort_value_name_idx_map_, FuncManager(), data_transfer_mgr_, op_kernel); diff --git a/onnxruntime/core/optimizer/optimizer_execution_frame.h b/onnxruntime/core/optimizer/optimizer_execution_frame.h index 944cf4fbe4..ba97bbd28a 100644 --- a/onnxruntime/core/optimizer/optimizer_execution_frame.h +++ b/onnxruntime/core/optimizer/optimizer_execution_frame.h @@ -21,14 +21,14 @@ class OptimizerExecutionFrame final : public IExecutionFrame { class Info { public: Info(const std::vector& nodes, const InitializedTensorSet& initialized_tensor_set, - std::unique_ptr cpu_execution_provider); + const IExecutionProvider& execution_provider); ~Info() { for (auto& kvp : deleter_for_initialized_tensors_) { kvp.second.f(kvp.second.param); } } AllocatorPtr GetAllocator(const OrtMemoryInfo& info) const { - return cpu_execution_provider_->GetAllocator(info.id, info.mem_type); + return execution_provider_.GetAllocator(info.id, info.mem_type); } AllocatorPtr GetAllocator() const { @@ -55,7 +55,6 @@ class OptimizerExecutionFrame final : public IExecutionFrame { private: // The optimizer is running on CPU execution provider by default. - std::unique_ptr cpu_execution_provider_; const int device_id_{0}; const OrtMemType mem_type_{OrtMemTypeDefault}; AllocatorPtr allocator_ptr_; @@ -69,6 +68,7 @@ class OptimizerExecutionFrame final : public IExecutionFrame { // munmap memory region and close file descriptor std::unordered_map deleter_for_initialized_tensors_; std::unique_ptr node_index_info_; + const IExecutionProvider& execution_provider_; ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(Info); }; diff --git a/onnxruntime/core/session/allocator_impl.h b/onnxruntime/core/session/allocator_impl.h index 5cfe100f89..4c0086af57 100644 --- a/onnxruntime/core/session/allocator_impl.h +++ b/onnxruntime/core/session/allocator_impl.h @@ -4,6 +4,8 @@ #pragma once #include "core/session/onnxruntime_c_api.h" #include "core/framework/allocator.h" +#include "core/framework/arena.h" +#include "core/session/device_allocator.h" namespace onnxruntime { class AllocatorWrapper : public IAllocator { @@ -19,4 +21,5 @@ class AllocatorWrapper : public IAllocator { private: OrtAllocator* impl_; }; + } // namespace onnxruntime diff --git a/onnxruntime/core/session/device_allocator.cc b/onnxruntime/core/session/device_allocator.cc index 61cae5a008..4717f3a864 100644 --- a/onnxruntime/core/session/device_allocator.cc +++ b/onnxruntime/core/session/device_allocator.cc @@ -1,47 +1,15 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include "core/framework/allocator.h" -#include "core/framework/utils.h" +#include "device_allocator.h" #include "core/session/inference_session.h" -#include "core/session/onnxruntime_cxx_api.h" -#include "core/session/ort_apis.h" -#include - -struct OrtAllocatorForDevice : public OrtAllocator { - - explicit OrtAllocatorForDevice(onnxruntime::AllocatorPtr&& dev_allocator) - : device_allocator_(std::move(dev_allocator)) { - OrtAllocator::version = ORT_API_VERSION; - OrtAllocator::Alloc = [](OrtAllocator* this_, size_t size) { return static_cast(this_)->Alloc(size); }; - OrtAllocator::Free = [](OrtAllocator* this_, void* p) { static_cast(this_)->Free(p); }; - OrtAllocator::Info = [](const OrtAllocator* this_) { return static_cast(this_)->Info(); }; - } - - ~OrtAllocatorForDevice() = default; - - void* Alloc(size_t size) const { - return device_allocator_->Alloc(size); - } - void Free(void* p) const { - device_allocator_->Free(p); - } - - const OrtMemoryInfo* Info() const { - return &device_allocator_->Info(); - } - - OrtAllocatorForDevice(const OrtAllocatorForDevice&) = delete; - OrtAllocatorForDevice& operator=(const OrtAllocatorForDevice&) = delete; - - private: - onnxruntime::AllocatorPtr device_allocator_; -}; +#include "core/session/ort_env.h" +#include "core/session/allocator_impl.h" #define API_IMPL_BEGIN try { #define API_IMPL_END \ } \ - catch (const std::exception& ex) { \ + catch (const std::exception& ex) { \ return OrtApis::CreateStatus(ORT_RUNTIME_EXCEPTION, ex.what()); \ } @@ -52,11 +20,82 @@ ORT_API_STATUS_IMPL(OrtApis::CreateAllocator, const OrtSession* sess, const OrtM if (!allocator_ptr) { return OrtApis::CreateStatus(ORT_INVALID_ARGUMENT, "No requested allocator available"); } - *out = new OrtAllocatorForDevice(std::move(allocator_ptr)); + *out = new onnxruntime::OrtAllocatorForDevice(std::move(allocator_ptr)); return nullptr; API_IMPL_END } -ORT_API(void, OrtApis::ReleaseAllocator, _Frees_ptr_opt_ OrtAllocator* allocator) { - delete reinterpret_cast(allocator); +ORT_API_STATUS_IMPL(OrtApis::CreateAndRegisterAllocator, _Inout_ OrtEnv* env, _In_ const OrtMemoryInfo* mem_info, + _In_ const OrtArenaCfg* arena_cfg) { + using namespace onnxruntime; + if (!env) { + return OrtApis::CreateStatus(ORT_INVALID_ARGUMENT, "Env is null"); + } + + if (!mem_info) { + return OrtApis::CreateStatus(ORT_INVALID_ARGUMENT, "OrtMemoryInfo is null"); + } + + // TODO should we allow sharing of non-CPU allocators? + if (mem_info->device.Type() != OrtDevice::CPU) { + return OrtApis::CreateStatus(ORT_INVALID_ARGUMENT, "Only CPU devices are supported for now."); + } + + // determine if arena should be used + bool create_arena = mem_info->alloc_type == OrtArenaAllocator; + +#ifdef USE_JEMALLOC +#if defined(USE_MIMALLOC_ARENA_ALLOCATOR) || defined(USE_MIMALLOC_STL_ALLOCATOR) +#error jemalloc and mimalloc should not both be enabled +#endif + //JEMalloc already has memory pool, so just use device allocator. + create_arena = false; +#elif !(defined(__amd64__) || defined(_M_AMD64)) + //Disable Arena allocator for x86_32 build because it may run into infinite loop when integer overflow happens + create_arena = false; +#endif + + AllocatorPtr allocator_ptr; + size_t max_mem = std::numeric_limits::max(); + + // create appropriate DeviceAllocatorRegistrationInfo and allocator based on create_arena + if (create_arena) { + ArenaExtendStrategy arena_extend_strategy = BFCArena::DEFAULT_ARENA_EXTEND_STRATEGY; + int initial_chunk_size_bytes = BFCArena::DEFAULT_INITIAL_CHUNK_SIZE_BYTES; + int max_dead_bytes_per_chunk = BFCArena::DEFAULT_MAX_DEAD_BYTES_PER_CHUNK; + if (arena_cfg) { + if (arena_cfg->max_mem != -1) max_mem = arena_cfg->max_mem; + if (arena_cfg->arena_extend_strategy == 0) { + arena_extend_strategy = ArenaExtendStrategy::kNextPowerOfTwo; + } else if (arena_cfg->arena_extend_strategy == 1) { + arena_extend_strategy = ArenaExtendStrategy::kSameAsRequested; + } + if (arena_cfg->initial_chunk_size_bytes != -1) initial_chunk_size_bytes = arena_cfg->initial_chunk_size_bytes; + if (arena_cfg->max_dead_bytes_per_chunk != -1) max_dead_bytes_per_chunk = arena_cfg->max_dead_bytes_per_chunk; + } + + DeviceAllocatorRegistrationInfo device_info{ + OrtMemTypeDefault, + [mem_info](int) { return onnxruntime::make_unique(*mem_info); }, + max_mem, + arena_extend_strategy, + initial_chunk_size_bytes, + max_dead_bytes_per_chunk}; + allocator_ptr = CreateAllocator(device_info, 0, create_arena); + } else { + DeviceAllocatorRegistrationInfo device_info{OrtMemTypeDefault, + [](int) { return onnxruntime::make_unique(); }, + max_mem}; + allocator_ptr = CreateAllocator(device_info, 0, create_arena); + } + + auto st = env->RegisterAllocator(allocator_ptr); + if (!st.IsOK()) { + return OrtApis::CreateStatus(ORT_INVALID_ARGUMENT, st.ErrorMessage().c_str()); + } + return nullptr; +} + +ORT_API(void, OrtApis::ReleaseAllocator, _Frees_ptr_opt_ OrtAllocator* allocator) { + delete reinterpret_cast(allocator); } diff --git a/onnxruntime/core/session/device_allocator.h b/onnxruntime/core/session/device_allocator.h new file mode 100644 index 0000000000..89e29b1ce1 --- /dev/null +++ b/onnxruntime/core/session/device_allocator.h @@ -0,0 +1,45 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +#pragma once + +#include "core/framework/allocator.h" +#include "core/framework/utils.h" +#include "core/session/onnxruntime_cxx_api.h" +#include "core/session/ort_apis.h" +#include +#include "core/framework/allocatormgr.h" + +namespace onnxruntime { +struct OrtAllocatorForDevice : public OrtAllocator { + explicit OrtAllocatorForDevice(onnxruntime::AllocatorPtr&& dev_allocator) + : device_allocator_(std::move(dev_allocator)) { + OrtAllocator::version = ORT_API_VERSION; + OrtAllocator::Alloc = [](OrtAllocator* this_, size_t size) { return static_cast(this_)->Alloc(size); }; + OrtAllocator::Free = [](OrtAllocator* this_, void* p) { static_cast(this_)->Free(p); }; + OrtAllocator::Info = [](const OrtAllocator* this_) { return static_cast(this_)->Info(); }; + } + + ~OrtAllocatorForDevice() = default; + + void* Alloc(size_t size) const { + return device_allocator_->Alloc(size); + } + void Free(void* p) const { + device_allocator_->Free(p); + } + + const OrtMemoryInfo* Info() const { + return &device_allocator_->Info(); + } + + OrtAllocatorForDevice(const OrtAllocatorForDevice&) = delete; + OrtAllocatorForDevice& operator=(const OrtAllocatorForDevice&) = delete; + + onnxruntime::IAllocator* GetAllocator() { + return device_allocator_.get(); + } + + private: + onnxruntime::AllocatorPtr device_allocator_; +}; +} // namespace onnxruntime diff --git a/onnxruntime/core/session/environment.cc b/onnxruntime/core/session/environment.cc index d8dead06d3..b8fbb162a3 100644 --- a/onnxruntime/core/session/environment.cc +++ b/onnxruntime/core/session/environment.cc @@ -20,6 +20,7 @@ #include "core/platform/env.h" #include "core/util/thread_utils.h" +#include "core/session/allocator_impl.h" #ifdef ONNXRUNTIME_ENABLE_INSTRUMENT #include "core/platform/tracing.h" @@ -48,6 +49,19 @@ Status Environment::Create(std::unique_ptr logging_mana return status; } +Status Environment::RegisterAllocator(AllocatorPtr allocator) { + const auto& mem_info = allocator->Info(); + // We don't expect millions of allocators getting registered. Hence linear search should be fine. + auto ite = std::find_if(std::begin(shared_allocators_), + std::end(shared_allocators_), + [&mem_info](const AllocatorPtr& alloc_ptr) { return alloc_ptr->Info() == mem_info; }); + if (ite != shared_allocators_.end()) { + return Status(ONNXRUNTIME, INVALID_ARGUMENT, "Allocator with this OrtMemoryInfo is already registered."); + } + shared_allocators_.insert(ite, allocator); + return Status::OK(); +} + Status Environment::Initialize(std::unique_ptr logging_manager, const OrtThreadingOptions* tp_options, bool create_global_thread_pools) { diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index f64d99f2d7..4df0e21648 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -233,7 +233,8 @@ InferenceSession::InferenceSession(const SessionOptions& session_options, const graph_transformation_mgr_(session_options.max_num_graph_transformation_steps), insert_cast_transformer_("CastFloat16Transformer"), #endif - logging_manager_(session_env.GetLoggingManager()) { + logging_manager_(session_env.GetLoggingManager()), + environment_(session_env) { // Initialize assets of this session instance ConstructorCommon(session_options, session_env); } @@ -244,7 +245,8 @@ InferenceSession::InferenceSession(const SessionOptions& session_options, const : model_location_(ToWideString(model_uri)), graph_transformation_mgr_(session_options.max_num_graph_transformation_steps), insert_cast_transformer_("CastFloat16Transformer"), - logging_manager_(session_env.GetLoggingManager()) { + logging_manager_(session_env.GetLoggingManager()), + environment_(session_env) { auto status = Model::Load(model_location_, model_proto_); ORT_ENFORCE(status.IsOK(), "Given model could not be parsed while creating inference session. Error message: ", status.ErrorMessage()); @@ -259,7 +261,8 @@ InferenceSession::InferenceSession(const SessionOptions& session_options, const std::wstring& model_uri) : graph_transformation_mgr_(session_options.max_num_graph_transformation_steps), insert_cast_transformer_("CastFloat16Transformer"), - logging_manager_(session_env.GetLoggingManager()) { + logging_manager_(session_env.GetLoggingManager()), + environment_(session_env) { model_location_ = ToWideString(model_uri); auto status = Model::Load(model_location_, model_proto_); ORT_ENFORCE(status.IsOK(), "Given model could not be parsed while creating inference session. Error message: ", @@ -274,7 +277,8 @@ InferenceSession::InferenceSession(const SessionOptions& session_options, const std::istream& model_istream) : graph_transformation_mgr_(session_options.max_num_graph_transformation_steps), insert_cast_transformer_("CastFloat16Transformer"), - logging_manager_(session_env.GetLoggingManager()) { + logging_manager_(session_env.GetLoggingManager()), + environment_(session_env) { Status st = Model::Load(model_istream, &model_proto_); ORT_ENFORCE(st.IsOK(), "Could not parse model successfully while constructing the inference session"); is_model_proto_parsed_ = true; @@ -286,7 +290,8 @@ InferenceSession::InferenceSession(const SessionOptions& session_options, const const void* model_data, int model_data_len) : graph_transformation_mgr_(session_options.max_num_graph_transformation_steps), insert_cast_transformer_("CastFloat16Transformer"), - logging_manager_(session_env.GetLoggingManager()) { + logging_manager_(session_env.GetLoggingManager()), + environment_(session_env) { const bool result = model_proto_.ParseFromArray(model_data, model_data_len); ORT_ENFORCE(result, "Could not parse model successfully while constructing the inference session"); is_model_proto_parsed_ = true; @@ -675,7 +680,7 @@ common::Status InferenceSession::TransformGraph(onnxruntime::Graph& graph, if (!node.Name().empty()) oss << node.Name() << ":"; oss << node.OpType() << "(" << node.SinceVersion() << ")"; - + return Status(common::ONNXRUNTIME, common::NOT_IMPLEMENTED, oss.str()); } else { if (is_verbose_mode) { // TODO: should we disable this if the number of nodes are above a certain threshold? @@ -804,6 +809,24 @@ common::Status InferenceSession::Initialize() { // re-acquire mutex std::lock_guard l(session_mutex_); + // At this time we know all the providers that will be part of this session. + // Read shared allocators from the environment and update them in the respective providers. + // + // The reason for updating the providers is so that when the session state is created the allocators + // are setup appropariately keyed by OrtMemoryInfo with delegates going to the respective providers. + // Secondly, the GetAllocator() method inside IExecutionProvider is still used in various places, hence + // it doesn't make sense to just update the allocator map inside session state with these shared allocators; doing + // so would cause inconsistency between the allocator map inside session sate and that inside the providers. + // TODO: we could refactor the allocators to not require the call to GetAllocator but that change is much bigger + // since we've to take into account the per-thread cuda allocators. + // TODO (contd.) We could also possibly absorb the per-thread logic in a new allocator decorator that derives + // from IAllocator to keep things clean. + std::string use_env_allocators = GetSessionConfigOrDefault(session_options_, + ORT_SESSION_OPTIONS_CONFIG_USE_ENV_ALLOCATORS, "0"); + if (use_env_allocators == "1") { + UpdateProvidersWithSharedAllocators(); + } + #ifdef ONNXRUNTIME_ENABLE_INSTRUMENT TraceLoggingWriteStart(session_activity, "OrtInferenceSessionActivity"); session_activity_started_ = true; @@ -914,6 +937,19 @@ common::Status InferenceSession::Initialize() { return status; } +// This method should be called from within Initialize() only and before the creation of the session state. +// This ensures all providers have been registered in the session and the session state is consistent with the providers. +void InferenceSession::UpdateProvidersWithSharedAllocators() { + using namespace std; + const auto& provider_ids = execution_providers_.GetIds(); + for (const auto& one_shared_alloc : environment_.GetRegisteredSharedAllocators()) { + for (const auto& id : provider_ids) { + auto* provider_ptr = execution_providers_.Get(id); + provider_ptr->ReplaceAllocator(one_shared_alloc); + } + } +} + int InferenceSession::GetCurrentNumRuns() const { return current_num_runs_.load(); } @@ -1473,7 +1509,9 @@ void InferenceSession::AddPredefinedTransformers(GraphTransformerManager& transf auto add_transformers = [&](TransformerLevel level) { // Generate and register transformers for level auto transformers_to_register = - optimizer_utils::GenerateTransformers(level, session_options_.free_dimension_overrides, custom_list); + optimizer_utils::GenerateTransformers(level, session_options_.free_dimension_overrides, + *execution_providers_.Get(onnxruntime::kCpuExecutionProvider), + custom_list); for (auto& entry : transformers_to_register) { transformer_manager.Register(std::move(entry), level); } diff --git a/onnxruntime/core/session/inference_session.h b/onnxruntime/core/session/inference_session.h index 8035ba6358..13c421dee7 100644 --- a/onnxruntime/core/session/inference_session.h +++ b/onnxruntime/core/session/inference_session.h @@ -335,6 +335,7 @@ class InferenceSession { * Get all the providers' options this session was initialized with. */ const ProviderOptionsMap& GetAllProviderOptions() const; + /** * Start profiling on this inference session. This simply turns on profiling events to be * recorded. A corresponding EndProfiling has to follow to write profiling data to a file. @@ -424,6 +425,9 @@ class InferenceSession { // The file path of where the model was loaded. e.g. /tmp/test_squeezenet/model.onnx std::basic_string model_location_; + // The list of execution providers. + ExecutionProviders execution_providers_; + private: ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(InferenceSession); @@ -467,6 +471,9 @@ class InferenceSession { template void StartProfiling(const std::basic_string& file_prefix); + // Updates all providers with the allocators from the env based on OrtMemoryInfo + void UpdateProvidersWithSharedAllocators(); + #if !defined(ORT_MINIMAL_BUILD) virtual void AddPredefinedTransformers(GraphTransformerManager& transformer_manager, TransformerLevel graph_optimization_level, @@ -499,9 +506,6 @@ class InferenceSession { // Profiler for this session. profiling::Profiler session_profiler_; - // The list of execution providers. - ExecutionProviders execution_providers_; - // Immutable state for each op in the model. Shared by all executors. // It has a dependency on execution_providers_. std::unique_ptr session_state_; @@ -588,6 +592,7 @@ class InferenceSession { // Flag indicating if ModelProto has been parsed in an applicable ctor bool is_model_proto_parsed_ = false; + const Environment& environment_; }; struct SessionIOBinding { diff --git a/onnxruntime/core/session/onnxruntime_c_api.cc b/onnxruntime/core/session/onnxruntime_c_api.cc index 89660149de..3bc74ed73f 100644 --- a/onnxruntime/core/session/onnxruntime_c_api.cc +++ b/onnxruntime/core/session/onnxruntime_c_api.cc @@ -1904,8 +1904,8 @@ static constexpr OrtApi ort_api_1_to_5 = { &OrtApis::GetBoundOutputValues, &OrtApis::ClearBoundInputs, &OrtApis::ClearBoundOutputs, - &OrtApis::TensorAt, + &OrtApis::CreateAndRegisterAllocator, }; // Assert to do a limited check to ensure Version 1 of OrtApi never changes (will detect an addition or deletion but not if they cancel out each other) diff --git a/onnxruntime/core/session/ort_apis.h b/onnxruntime/core/session/ort_apis.h index 7f548dd724..9b8ed5be3c 100644 --- a/onnxruntime/core/session/ort_apis.h +++ b/onnxruntime/core/session/ort_apis.h @@ -230,4 +230,5 @@ ORT_API_STATUS_IMPL(AddSessionConfigEntry, _Inout_ OrtSessionOptions* options, ORT_API_STATUS_IMPL(TensorAt, _Inout_ OrtValue* value, size_t* location_values, size_t location_values_count, _Outptr_ void** out); +ORT_API_STATUS_IMPL(CreateAndRegisterAllocator, _Inout_ OrtEnv* env, _In_ const OrtMemoryInfo* mem_info, _In_ const OrtArenaCfg* arena_cfg); } // namespace OrtApis diff --git a/onnxruntime/core/session/ort_env.cc b/onnxruntime/core/session/ort_env.cc index 7bcc10d924..f7f460eec9 100644 --- a/onnxruntime/core/session/ort_env.cc +++ b/onnxruntime/core/session/ort_env.cc @@ -8,6 +8,7 @@ #include "ort_env.h" #include "core/session/ort_apis.h" #include "core/session/environment.h" +#include "core/session/allocator_impl.h" #include "core/common/logging/logging.h" #ifdef __ANDROID__ #include "core/platform/android/logging/android_log_sink.h" @@ -100,4 +101,9 @@ onnxruntime::logging::LoggingManager* OrtEnv::GetLoggingManager() const { void OrtEnv::SetLoggingManager(std::unique_ptr logging_manager) { value_->SetLoggingManager(std::move(logging_manager)); +} + +onnxruntime::Status OrtEnv::RegisterAllocator(AllocatorPtr allocator) { + auto status = value_->RegisterAllocator(allocator); + return status; } \ No newline at end of file diff --git a/onnxruntime/core/session/ort_env.h b/onnxruntime/core/session/ort_env.h index 7ecb603248..041b072a93 100644 --- a/onnxruntime/core/session/ort_env.h +++ b/onnxruntime/core/session/ort_env.h @@ -8,6 +8,7 @@ #include "core/common/logging/isink.h" #include "core/platform/ort_mutex.h" #include "core/common/status.h" +#include "core/framework/allocator.h" namespace onnxruntime { class Environment; @@ -55,6 +56,12 @@ struct OrtEnv { onnxruntime::logging::LoggingManager* GetLoggingManager() const; void SetLoggingManager(std::unique_ptr logging_manager); + /** + * Registers an allocator for sharing between multiple sessions. + * Returns an error if an allocator with the same OrtMemoryInfo is already registered. + */ + onnxruntime::Status RegisterAllocator(onnxruntime::AllocatorPtr allocator); + private: static OrtEnv* p_instance_; static onnxruntime::OrtMutex m_; diff --git a/onnxruntime/test/framework/inference_session_test.cc b/onnxruntime/test/framework/inference_session_test.cc index 609ca3b441..ab7f30a443 100644 --- a/onnxruntime/test/framework/inference_session_test.cc +++ b/onnxruntime/test/framework/inference_session_test.cc @@ -33,6 +33,8 @@ #endif #include "core/session/environment.h" #include "core/session/IOBinding.h" +#include "core/session/device_allocator.h" +#include "core/session/allocator_impl.h" #include "dummy_provider.h" #include "test_utils.h" #include "test/capturing_sink.h" @@ -2387,5 +2389,109 @@ TEST(InferenceSessionTests, InvalidSessionEnvCombination) { } } +// Tests for sharing allocators between sessions +class InferenceSessionTestSharingAllocator : public InferenceSession { + public: + InferenceSessionTestSharingAllocator(const SessionOptions& session_options, + const Environment& env) + : InferenceSession(session_options, env) { + } + + const SessionState& GetSessionState() { return InferenceSession::GetSessionState(); } +}; + +// Ensure sessions use the same allocator. It uses ORT created allocator. +TEST(InferenceSessionTests, AllocatorSharing_EnsureSessionsUseSameOrtCreatedAllocator) { + auto logging_manager = onnxruntime::make_unique( + std::unique_ptr(new CLogSink()), logging::Severity::kVERBOSE, false, + LoggingManager::InstanceType::Temporal); + + std::unique_ptr env; + auto st = Environment::Create(std::move(logging_manager), env); + ASSERT_TRUE(st.IsOK()); + // create allocator to register with the env + bool use_arena = true; +#if !(defined(__amd64__) || defined(_M_AMD64)) + use_arena = false; +#endif + OrtMemoryInfo mem_info{onnxruntime::CPU, use_arena ? OrtArenaAllocator : OrtDeviceAllocator}; + size_t max_mem = std::numeric_limits::max(); + DeviceAllocatorRegistrationInfo device_info{ + OrtMemTypeDefault, + [mem_info](int) { return onnxruntime::make_unique(mem_info); }, + max_mem}; + + AllocatorPtr allocator_ptr = CreateAllocator(device_info, 0, use_arena); + st = env->RegisterAllocator(allocator_ptr); + ASSERT_STATUS_OK(st); + // create sessions to share the allocator + + SessionOptions so1; + AddSessionConfigEntryImpl(so1, ORT_SESSION_OPTIONS_CONFIG_USE_ENV_ALLOCATORS, "1"); + InferenceSessionTestSharingAllocator sess1(so1, *env); + ASSERT_STATUS_OK(sess1.Load(MODEL_URI)); + ASSERT_STATUS_OK(sess1.Initialize()); + + SessionOptions so2; + AddSessionConfigEntryImpl(so2, ORT_SESSION_OPTIONS_CONFIG_USE_ENV_ALLOCATORS, "1"); + InferenceSessionTestSharingAllocator sess2(so2, *env); + ASSERT_STATUS_OK(sess2.Load(MODEL_URI)); + ASSERT_STATUS_OK(sess2.Initialize()); + + // This line ensures the allocator in the session is the same as that in the env + ASSERT_EQ(sess1.GetSessionState().GetAllocator(mem_info).get(), + allocator_ptr.get()); + + // This line ensures the underlying IAllocator* is the same across 2 sessions. + ASSERT_EQ(sess1.GetSessionState().GetAllocator(mem_info).get(), + sess2.GetSessionState().GetAllocator(mem_info).get()); +} + +// Ensure sessions don't use the same allocator. It uses ORT created allocator. +TEST(InferenceSessionTests, AllocatorSharing_EnsureSessionsDontUseSameOrtCreatedAllocator) { + auto logging_manager = onnxruntime::make_unique( + std::unique_ptr(new CLogSink()), logging::Severity::kVERBOSE, false, + LoggingManager::InstanceType::Temporal); + + std::unique_ptr env; + auto st = Environment::Create(std::move(logging_manager), env); + ASSERT_TRUE(st.IsOK()); + // create allocator to register with the env + bool use_arena = true; +#if !(defined(__amd64__) || defined(_M_AMD64)) + use_arena = false; +#endif + OrtMemoryInfo mem_info{onnxruntime::CPU, use_arena ? OrtArenaAllocator : OrtDeviceAllocator}; + size_t max_mem = std::numeric_limits::max(); + DeviceAllocatorRegistrationInfo device_info{ + OrtMemTypeDefault, + [mem_info](int) { return onnxruntime::make_unique(mem_info); }, + max_mem}; + + AllocatorPtr allocator_ptr = CreateAllocator(device_info, 0, use_arena); + st = env->RegisterAllocator(allocator_ptr); + ASSERT_STATUS_OK(st); + // create sessions to share the allocator + + SessionOptions so1; + AddSessionConfigEntryImpl(so1, ORT_SESSION_OPTIONS_CONFIG_USE_ENV_ALLOCATORS, "1"); + InferenceSessionTestSharingAllocator sess1(so1, *env); + ASSERT_STATUS_OK(sess1.Load(MODEL_URI)); + ASSERT_STATUS_OK(sess1.Initialize()); + + SessionOptions so2; + AddSessionConfigEntryImpl(so2, ORT_SESSION_OPTIONS_CONFIG_USE_ENV_ALLOCATORS, "0"); + InferenceSessionTestSharingAllocator sess2(so2, *env); + ASSERT_STATUS_OK(sess2.Load(MODEL_URI)); + ASSERT_STATUS_OK(sess2.Initialize()); + + // This line ensures the allocator in the session is the same as that in the env + ASSERT_EQ(sess1.GetSessionState().GetAllocator(mem_info).get(), + allocator_ptr.get()); + + // This line ensures the underlying OrtAllocator* is the same across 2 sessions. + ASSERT_NE(sess1.GetSessionState().GetAllocator(mem_info).get(), + sess2.GetSessionState().GetAllocator(mem_info).get()); +} } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/optimizer/cse_test.cc b/onnxruntime/test/optimizer/cse_test.cc index f0c7b24de7..71ed34fe51 100644 --- a/onnxruntime/test/optimizer/cse_test.cc +++ b/onnxruntime/test/optimizer/cse_test.cc @@ -231,10 +231,11 @@ TEST(CseTests, MergeConstants) { GraphTransformerManager graph_transformation_mgr(1); // In current implementation, equal constants are not merged. So CSE must precede constant folding, otherwise we end up // with multiple copies of the same constant. + std::unique_ptr e = onnxruntime::make_unique(CPUExecutionProviderInfo()); ASSERT_TRUE( graph_transformation_mgr.Register(onnxruntime::make_unique(), TransformerLevel::Level1).IsOK()); ASSERT_TRUE( - graph_transformation_mgr.Register(onnxruntime::make_unique(), TransformerLevel::Level1).IsOK()); + graph_transformation_mgr.Register(onnxruntime::make_unique(*e.get()), TransformerLevel::Level1).IsOK()); ASSERT_TRUE( graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, DefaultLoggingManager().DefaultLogger()).IsOK()); diff --git a/onnxruntime/test/optimizer/graph_transform_test.cc b/onnxruntime/test/optimizer/graph_transform_test.cc index 9a7b4fef57..6d7cb51abc 100644 --- a/onnxruntime/test/optimizer/graph_transform_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_test.cc @@ -139,9 +139,10 @@ TEST_F(GraphTransformationTests, ConstantFolding) { Graph& graph = model->MainGraph(); std::map op_to_count = CountOpsInGraph(graph); ASSERT_TRUE(op_to_count["Unsqueeze"] == 2); - + std::unique_ptr e = + onnxruntime::make_unique(CPUExecutionProviderInfo()); onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; - graph_transformation_mgr.Register(onnxruntime::make_unique(), TransformerLevel::Level1); + graph_transformation_mgr.Register(onnxruntime::make_unique(*e.get()), TransformerLevel::Level1); ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); @@ -156,9 +157,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingNodesOnDifferentEP) { Graph& graph = model->MainGraph(); std::map op_to_count = CountOpsInGraph(graph); ASSERT_TRUE(op_to_count["Unsqueeze"] == 2); - + std::unique_ptr e = + onnxruntime::make_unique(CPUExecutionProviderInfo()); onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; - graph_transformation_mgr.Register(onnxruntime::make_unique(), TransformerLevel::Level1); + graph_transformation_mgr.Register(onnxruntime::make_unique(*e.get()), TransformerLevel::Level1); // assign all nodes to CUDA. the constant folding should override this to perform the constant folding on cpu for (auto& node : graph.Nodes()) { @@ -236,9 +238,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingSubgraph) { std::map op_to_count = CountOpsInGraph(graph); ASSERT_TRUE(op_to_count["Add"] == 2); // one in each subgraph - + std::unique_ptr e = + onnxruntime::make_unique(CPUExecutionProviderInfo()); onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; - graph_transformation_mgr.Register(onnxruntime::make_unique(), TransformerLevel::Level1); + graph_transformation_mgr.Register(onnxruntime::make_unique(*e.get()), TransformerLevel::Level1); ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); @@ -261,7 +264,9 @@ TEST_F(GraphTransformationTests, ConstantFoldingWithShapeToInitializer) { std::unordered_set excluded_initializers; excluded_initializers.insert("matmul_weight"); onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; - graph_transformation_mgr.Register(onnxruntime::make_unique(compatible_eps, excluded_initializers), TransformerLevel::Level1); + std::unique_ptr e = + onnxruntime::make_unique(CPUExecutionProviderInfo()); + graph_transformation_mgr.Register(onnxruntime::make_unique(*e.get(), compatible_eps, excluded_initializers), TransformerLevel::Level1); ASSERT_TRUE(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_).IsOK()); @@ -283,7 +288,9 @@ TEST_F(GraphTransformationTests, ConstantFoldingWithScalarShapeToInitializer) { std::unordered_set compatible_eps; onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; - graph_transformation_mgr.Register(onnxruntime::make_unique(compatible_eps), TransformerLevel::Level1); + std::unique_ptr e = + onnxruntime::make_unique(CPUExecutionProviderInfo()); + graph_transformation_mgr.Register(onnxruntime::make_unique(*e.get(), compatible_eps), TransformerLevel::Level1); ASSERT_TRUE(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_).IsOK()); diff --git a/onnxruntime/test/optimizer/graph_transform_utils_test.cc b/onnxruntime/test/optimizer/graph_transform_utils_test.cc index e4b15e1f10..92e36a455c 100644 --- a/onnxruntime/test/optimizer/graph_transform_utils_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_utils_test.cc @@ -41,8 +41,9 @@ TEST(GraphTransformerUtilsTests, TestGenerateGraphTransformers) { std::string l1_transformer = "ConstantFolding"; std::string l2_transformer = "ConvActivationFusion"; std::vector custom_list = {l1_rule1, l1_transformer, l2_transformer}; - - auto transformers = optimizer_utils::GenerateTransformers(TransformerLevel::Level1, {}, custom_list); + std::unique_ptr cpu_execution_provider = + onnxruntime::make_unique(CPUExecutionProviderInfo()); + auto transformers = optimizer_utils::GenerateTransformers(TransformerLevel::Level1, {}, *cpu_execution_provider.get(), custom_list); ASSERT_TRUE(transformers.size() == 2); auto l1_rule_transformer_name = optimizer_utils::GenerateRuleBasedTransformerName(TransformerLevel::Level1); @@ -54,7 +55,7 @@ TEST(GraphTransformerUtilsTests, TestGenerateGraphTransformers) { } ASSERT_TRUE(rule_transformer && rule_transformer->RulesCount() == 1); - transformers = optimizer_utils::GenerateTransformers(TransformerLevel::Level2, {}, custom_list); + transformers = optimizer_utils::GenerateTransformers(TransformerLevel::Level2, {}, *cpu_execution_provider.get(), custom_list); #ifndef DISABLE_CONTRIB_OPS ASSERT_TRUE(transformers.size() == 1); #else @@ -65,15 +66,17 @@ TEST(GraphTransformerUtilsTests, TestGenerateGraphTransformers) { TEST(GraphTransformerUtilsTests, TestCustomOnlyTransformers) { // Transformers that are disabled by default. They can only be enabled by custom list. std::string l2_transformer = "GeluApproximation"; + std::unique_ptr cpu_execution_provider = + onnxruntime::make_unique(CPUExecutionProviderInfo()); std::vector default_list = {}; - auto default_transformers = optimizer_utils::GenerateTransformers(TransformerLevel::Level2, {}, default_list); + auto default_transformers = optimizer_utils::GenerateTransformers(TransformerLevel::Level2, {}, *cpu_execution_provider.get(), default_list); for (auto& transformer : default_transformers) { ASSERT_TRUE(transformer->Name() != l2_transformer); } std::vector custom_list = {l2_transformer}; - auto custom_transformers = optimizer_utils::GenerateTransformers(TransformerLevel::Level2, {}, custom_list); + auto custom_transformers = optimizer_utils::GenerateTransformers(TransformerLevel::Level2, {}, *cpu_execution_provider.get(), custom_list); #ifndef DISABLE_CONTRIB_OPS ASSERT_TRUE(custom_transformers.size() == 1); ASSERT_TRUE(custom_transformers[0]->Name() == l2_transformer); diff --git a/onnxruntime/test/optimizer/optimizer_test.cc b/onnxruntime/test/optimizer/optimizer_test.cc index af550b3dc4..8b76fb2fc3 100644 --- a/onnxruntime/test/optimizer/optimizer_test.cc +++ b/onnxruntime/test/optimizer/optimizer_test.cc @@ -66,7 +66,7 @@ TEST(OptimizerTest, Basic) { std::unique_ptr cpu_execution_provider = onnxruntime::make_unique(CPUExecutionProviderInfo()); - OptimizerExecutionFrame::Info info(nodes, initialized_tensor_set, std::move(cpu_execution_provider)); + OptimizerExecutionFrame::Info info(nodes, initialized_tensor_set, *cpu_execution_provider.get()); std::vector fetch_mlvalue_idxs{info.GetMLValueIndex("out")}; OptimizerExecutionFrame frame(info, fetch_mlvalue_idxs); const logging::Logger& logger = DefaultLoggingManager().DefaultLogger(); diff --git a/onnxruntime/test/shared_lib/test_inference.cc b/onnxruntime/test/shared_lib/test_inference.cc index 3b5f169d03..839b942cb0 100644 --- a/onnxruntime/test/shared_lib/test_inference.cc +++ b/onnxruntime/test/shared_lib/test_inference.cc @@ -826,3 +826,58 @@ TEST(CApiTest, get_available_providers_cpp) { ASSERT_TRUE(providers.size() > 0); ASSERT_TRUE(providers[0] == std::string("CPUExecutionProvider")); } + +// This test uses the CreateAndRegisterAllocator API to register an allocator with the env, +// creates 2 sessions and then runs those 2 sessions one after another +TEST(CApiTest, TestSharedAllocatorUsingCreateAndRegisterAllocator) { + // simple inference test + // prepare inputs + std::vector inputs(1); + Input& input = inputs.back(); + input.name = "X"; + input.dims = {3, 2}; + input.values = {1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f}; + + // prepare expected inputs and outputs + std::vector expected_dims_y = {3, 2}; + std::vector expected_values_y = {1.0f, 4.0f, 9.0f, 16.0f, 25.0f, 36.0f}; + OrtEnv* env_ptr = (OrtEnv*)(*ort_env); + + OrtMemoryInfo* mem_info = nullptr; + const auto& api = Ort::GetApi(); + std::unique_ptr rel_info(mem_info, api.ReleaseMemoryInfo); + ASSERT_TRUE(api.CreateCpuMemoryInfo(OrtArenaAllocator, OrtMemTypeDefault, &mem_info) == nullptr); + + OrtArenaCfg arena_cfg{-1, -1, -1, -1}; + ASSERT_TRUE(api.CreateAndRegisterAllocator(env_ptr, mem_info, &arena_cfg) == nullptr); + + // test for duplicates + std::unique_ptr status_releaser(api.CreateAndRegisterAllocator(env_ptr, mem_info, + &arena_cfg), + api.ReleaseStatus); + ASSERT_FALSE(status_releaser.get() == nullptr); + + Ort::SessionOptions session_options; + auto default_allocator = onnxruntime::make_unique(); + session_options.AddConfigEntry(ORT_SESSION_OPTIONS_CONFIG_USE_ENV_ALLOCATORS, "1"); + + // create session 1 + Ort::Session session1(*ort_env, MODEL_URI, session_options); + RunSession(default_allocator.get(), + session1, + inputs, + "Y", + expected_dims_y, + expected_values_y, + nullptr); + + // create session 2 + Ort::Session session2(*ort_env, MODEL_URI, session_options); + RunSession(default_allocator.get(), + session2, + inputs, + "Y", + expected_dims_y, + expected_values_y, + nullptr); +} \ No newline at end of file diff --git a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc index 7dc99e7711..91b6d37846 100644 --- a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc +++ b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc @@ -51,6 +51,7 @@ std::vector> GeneratePreTrainingTransformers( TransformerLevel level, const std::unordered_set& weights_to_train, const TrainingSession::TrainingConfiguration::GraphTransformerConfiguration& config, + const IExecutionProvider& execution_provider, const std::vector& transformers_and_rules_to_enable) { std::vector> transformers; std::unique_ptr rule_transformer = nullptr; @@ -88,7 +89,7 @@ std::vector> GeneratePreTrainingTransformers( transformers.emplace_back(onnxruntime::make_unique(compatible_eps)); } - transformers.emplace_back(onnxruntime::make_unique(compatible_eps, weights_to_train)); + transformers.emplace_back(onnxruntime::make_unique(execution_provider, compatible_eps, weights_to_train)); auto horizontal_parallel_size = training::DistributedRunContext::GroupSize(training::WorkerGroupType::HorizontalParallel); if (horizontal_parallel_size > 1) { LOGS_DEFAULT(WARNING) << horizontal_parallel_size << "-way horizontal model parallel is enabled"; diff --git a/orttraining/orttraining/core/optimizer/graph_transformer_utils.h b/orttraining/orttraining/core/optimizer/graph_transformer_utils.h index c7b6cfe085..d6f932742e 100644 --- a/orttraining/orttraining/core/optimizer/graph_transformer_utils.h +++ b/orttraining/orttraining/core/optimizer/graph_transformer_utils.h @@ -19,6 +19,7 @@ std::vector> GeneratePreTrainingTransformers( TransformerLevel level, const std::unordered_set& weights_to_train, const TrainingSession::TrainingConfiguration::GraphTransformerConfiguration& config, + const IExecutionProvider& execution_provider, // required for constant folding const std::vector& rules_and_transformers_to_enable = {}); /** Generates all predefined (both rule-based and non-rule-based) transformers for this level. diff --git a/orttraining/orttraining/core/session/training_session.cc b/orttraining/orttraining/core/session/training_session.cc index a8dcf90b06..c4e628ecbc 100644 --- a/orttraining/orttraining/core/session/training_session.cc +++ b/orttraining/orttraining/core/session/training_session.cc @@ -7,6 +7,7 @@ #include "core/graph/model.h" #include "core/session/IOBinding.h" #include "core/providers/cpu/controlflow/utils.h" +#include "core/providers/cpu/cpu_execution_provider.h" #include "orttraining/core/graph/loss_function_builder.h" #include "orttraining/core/graph/optimizer_builder.h" #include "orttraining/core/framework/checkpointing.h" @@ -492,7 +493,14 @@ static Status AddGradientAccumulationNodes(Graph& graph, Status TrainingSession::ApplyTransformationsToMainGraph(const std::unordered_set& weights_to_train, const TrainingConfiguration::GraphTransformerConfiguration& config) { GraphTransformerManager graph_transformation_mgr{1}; - AddPreTrainingTransformers(graph_transformation_mgr, weights_to_train, config); + // TODO: ideally we can just reuse the CPU EP registered with the session, but in the training session case + // the EPs are registered after ConfigureForTraining and before Initialize is called. Hence we don't have access + // to the registered CPU EP at this stage. Hence creating the EP here again. This is still much better than + // creating an EP instance for every single node in ConstantFolding. + // Create execution frame for executing constant nodes. + std::unique_ptr cpu_execution_provider = + onnxruntime::make_unique(CPUExecutionProviderInfo()); + AddPreTrainingTransformers(*cpu_execution_provider, graph_transformation_mgr, weights_to_train, config); // apply transformers Graph& graph = model_->MainGraph(); @@ -504,15 +512,17 @@ Status TrainingSession::ApplyTransformationsToMainGraph(const std::unordered_set } // Registers all the pre transformers with transformer manager -void TrainingSession::AddPreTrainingTransformers(GraphTransformerManager& transformer_manager, +void TrainingSession::AddPreTrainingTransformers(const IExecutionProvider& execution_provider, + GraphTransformerManager& transformer_manager, const std::unordered_set& weights_to_train, const TrainingConfiguration::GraphTransformerConfiguration& config, TransformerLevel graph_optimization_level, const std::vector& custom_list) { auto add_transformers = [&](TransformerLevel level) { // Generate and register transformers for level + auto transformers_to_register = transformer_utils::GeneratePreTrainingTransformers( - level, weights_to_train, config, custom_list); + level, weights_to_train, config, execution_provider, custom_list); for (auto& entry : transformers_to_register) { transformer_manager.Register(std::move(entry), level); } diff --git a/orttraining/orttraining/core/session/training_session.h b/orttraining/orttraining/core/session/training_session.h index b643675607..a8d7c0b4c9 100644 --- a/orttraining/orttraining/core/session/training_session.h +++ b/orttraining/orttraining/core/session/training_session.h @@ -400,7 +400,8 @@ class TrainingSession : public InferenceSession { const TrainingConfiguration::GraphTransformerConfiguration& config); /** configure initial transformers for training */ - void AddPreTrainingTransformers(GraphTransformerManager& transformer_manager, + void AddPreTrainingTransformers(const IExecutionProvider& execution_provider, // for constant folding + GraphTransformerManager& transformer_manager, const std::unordered_set& weights_to_train, const TrainingConfiguration::GraphTransformerConfiguration& config, TransformerLevel graph_optimization_level = TransformerLevel::MaxLevel, From db7669b22576309a7d954a1232dff73c503d8e3f Mon Sep 17 00:00:00 2001 From: Scott McKay Date: Sun, 23 Aug 2020 07:02:13 +1000 Subject: [PATCH 2/4] Reduce ONNX dependency in minimal build (#4890) * Next round of changes. Remove inclusion of ONNX schema header Exclude custom registry related things Move IsConstantInitializer from graph_utils to Graph as it's needed in a minimal build and graph_utils is excluded. --- cmake/onnxruntime_framework.cmake | 2 + cmake/onnxruntime_graph.cmake | 5 ++- .../onnxruntime/core/framework/op_kernel.h | 2 +- include/onnxruntime/core/graph/basic_types.h | 5 +++ include/onnxruntime/core/graph/graph.h | 7 ++++ .../onnxruntime/core/graph/onnx_protobuf.h | 37 ++++++++++++------- onnxruntime/contrib_ops/cpu/expand_dims.cc | 4 +- onnxruntime/contrib_ops/cpu/sample.cc | 3 +- onnxruntime/contrib_ops/cpu/tokenizer.cc | 5 +-- onnxruntime/core/framework/kernel_registry.cc | 12 +++++- .../core/framework/kernel_registry_manager.cc | 5 ++- .../core/framework/kernel_registry_manager.h | 1 - .../core/framework/op_node_proto_helper.cc | 23 ++++++++---- .../core/framework/provider_bridge_ort.cc | 21 ++++++----- .../core/graph/contrib_ops/contrib_defs.cc | 2 +- .../graph/featurizers_ops/featurizers_defs.cc | 4 +- onnxruntime/core/graph/graph.cc | 29 ++++++++++++++- onnxruntime/core/graph/graph_utils.cc | 22 +---------- onnxruntime/core/graph/graph_viewer.cc | 4 +- onnxruntime/core/graph/model.cc | 3 ++ .../providers/cpu/nn/string_normalizer.cc | 12 +++--- .../core/providers/cpu/nn/tfidfvectorizer.cc | 1 - .../providers/cpu/tensor/reverse_sequence.cc | 1 - onnxruntime/core/session/custom_ops.cc | 6 ++- onnxruntime/core/session/environment.cc | 4 ++ onnxruntime/core/session/inference_session.cc | 33 +++++++++-------- .../python/onnxruntime_pybind_state.cc | 13 ++++--- .../test/framework/allocation_planner_test.cc | 9 +++-- onnxruntime/test/framework/allocator_test.cc | 1 + .../test/framework/session_state_test.cc | 3 +- 30 files changed, 170 insertions(+), 109 deletions(-) diff --git a/cmake/onnxruntime_framework.cmake b/cmake/onnxruntime_framework.cmake index a3ae8f4fdd..d8697b9c47 100644 --- a/cmake/onnxruntime_framework.cmake +++ b/cmake/onnxruntime_framework.cmake @@ -11,6 +11,8 @@ if (onnxruntime_MINIMAL_BUILD) file(GLOB onnxruntime_framework_src_exclude "${ONNXRUNTIME_ROOT}/core/framework/provider_bridge_ort.cc" "${ONNXRUNTIME_ROOT}/core/framework/graph_partitioner.*" + "${ONNXRUNTIME_INCLUDE_DIR}/core/framework/customregistry.h" + "${ONNXRUNTIME_ROOT}/core/framework/customregistry.cc" ) list(REMOVE_ITEM onnxruntime_framework_srcs ${onnxruntime_framework_src_exclude}) diff --git a/cmake/onnxruntime_graph.cmake b/cmake/onnxruntime_graph.cmake index e134db0e86..dadc3771f0 100644 --- a/cmake/onnxruntime_graph.cmake +++ b/cmake/onnxruntime_graph.cmake @@ -11,10 +11,13 @@ file(GLOB_RECURSE onnxruntime_graph_src CONFIGURE_DEPENDS set(onnxruntime_graph_src_exclude_patterns) if (onnxruntime_MINIMAL_BUILD) - # remove schema registration of contrib ops + # remove schema registration support list(APPEND onnxruntime_graph_src_exclude_patterns + "${ONNXRUNTIME_INCLUDE_DIR}/core/graph/schema_registry.h" + "${ONNXRUNTIME_ROOT}/core/graph/schema_registry.cc" "${ONNXRUNTIME_ROOT}/core/graph/contrib_ops/*defs.h" "${ONNXRUNTIME_ROOT}/core/graph/contrib_ops/*defs.cc" + ) # no Function support initially diff --git a/include/onnxruntime/core/framework/op_kernel.h b/include/onnxruntime/core/framework/op_kernel.h index e26d2f62a7..835deb027a 100644 --- a/include/onnxruntime/core/framework/op_kernel.h +++ b/include/onnxruntime/core/framework/op_kernel.h @@ -17,8 +17,8 @@ #include "core/framework/sparse_tensor.h" #include "core/graph/constants.h" #include "core/graph/graph_viewer.h" +#include "core/graph/onnx_protobuf.h" #include "gsl/gsl" -#include "onnx/defs/schema.h" namespace onnxruntime { class IExecutionFrame; diff --git a/include/onnxruntime/core/graph/basic_types.h b/include/onnxruntime/core/graph/basic_types.h index 24702b4df7..a0085a5ad9 100644 --- a/include/onnxruntime/core/graph/basic_types.h +++ b/include/onnxruntime/core/graph/basic_types.h @@ -12,6 +12,11 @@ class ValueInfoProto; class TensorProto; class TypeProto; class AttributeProto; +// define types that would come from the ONNX library if we were building against it. +#if defined(ORT_MINIMAL_BUILD) +using OperatorSetVersion = int; +#endif + } // namespace ONNX_NAMESPACE namespace onnxruntime { diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index ced837e7e7..1f4d2fd5bc 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -578,6 +578,13 @@ class Graph { /** Returns true if an initializer value can be overridden by a graph input with the same name. */ bool CanOverrideInitializer() const noexcept { return ir_version_ >= 4; } + /** returns the initializer's TensorProto if 'name' is an initializer, is constant and + cannot be overridden at runtime. If the initializer is not found or is not constant, a nullptr is returned. + @param check_outer_scope If true and the graph is a subgraph, + check ancestor graph/s for 'name' if not found in 'graph'. + */ + const ONNX_NAMESPACE::TensorProto* GetConstantInitializer(const std::string& name, bool check_outer_scope) const; + /** Gets the Graph inputs excluding initializers. These are the required inputs to the Graph as the initializers can be optionally overridden via graph inputs. @remarks Contains no nullptr values. */ diff --git a/include/onnxruntime/core/graph/onnx_protobuf.h b/include/onnxruntime/core/graph/onnx_protobuf.h index 40d55a9581..755edf9784 100644 --- a/include/onnxruntime/core/graph/onnx_protobuf.h +++ b/include/onnxruntime/core/graph/onnx_protobuf.h @@ -18,22 +18,33 @@ #pragma warning(disable : 4100) #pragma warning(disable : 4146) /*unary minus operator applied to unsigned type, result still unsigned*/ #pragma warning(disable : 4127) -#pragma warning(disable : 4244) /*'conversion' conversion from 'type1' to 'type2', possible loss of data*/ -#pragma warning(disable : 4251) /*'identifier' : class 'type' needs to have dll-interface to be used by clients of class 'type2'*/ -#pragma warning(disable : 4267) /*'var' : conversion from 'size_t' to 'type', possible loss of data*/ -#pragma warning(disable : 4305) /*'identifier' : truncation from 'type1' to 'type2'*/ -#pragma warning(disable : 4307) /*'operator' : integral constant overflow*/ -#pragma warning(disable : 4309) /*'conversion' : truncation of constant value*/ -#pragma warning(disable : 4334) /*'operator' : result of 32-bit shift implicitly converted to 64 bits (was 64-bit shift intended?)*/ -#pragma warning(disable : 4355) /*'this' : used in base member initializer list*/ -#pragma warning(disable : 4506) /*no definition for inline function 'function'*/ -#pragma warning(disable : 4800) /*'type' : forcing value to bool 'true' or 'false' (performance warning)*/ -#pragma warning(disable : 4996) /*The compiler encountered a deprecated declaration.*/ -#pragma warning(disable : 6011) /*Dereferencing NULL pointer*/ -#pragma warning(disable : 6387) /*'value' could be '0'*/ +#pragma warning(disable : 4244) /*'conversion' conversion from 'type1' to 'type2', possible loss of data*/ +#pragma warning(disable : 4251) /*'identifier' : class 'type' needs to have dll-interface to be used by clients of class 'type2'*/ +#pragma warning(disable : 4267) /*'var' : conversion from 'size_t' to 'type', possible loss of data*/ +#pragma warning(disable : 4305) /*'identifier' : truncation from 'type1' to 'type2'*/ +#pragma warning(disable : 4307) /*'operator' : integral constant overflow*/ +#pragma warning(disable : 4309) /*'conversion' : truncation of constant value*/ +#pragma warning(disable : 4334) /*'operator' : result of 32-bit shift implicitly converted to 64 bits (was 64-bit shift intended?)*/ +#pragma warning(disable : 4355) /*'this' : used in base member initializer list*/ +#pragma warning(disable : 4506) /*no definition for inline function 'function'*/ +#pragma warning(disable : 4800) /*'type' : forcing value to bool 'true' or 'false' (performance warning)*/ +#pragma warning(disable : 4996) /*The compiler encountered a deprecated declaration.*/ +#pragma warning(disable : 6011) /*Dereferencing NULL pointer*/ +#pragma warning(disable : 6387) /*'value' could be '0'*/ #pragma warning(disable : 26495) /*Variable is uninitialized.*/ #endif + +#if !defined(ORT_MINIMAL_BUILD) #include "onnx/defs/schema.h" +#else +#include "onnx/defs/data_type_utils.h" + +// stub definition of OpSchema to minimize other code changes +namespace ONNX_NAMESPACE { +class OpSchema {}; +} // namespace ONNX_NAMESPACE +#endif + #include "onnx/onnx_pb.h" #ifdef __GNUC__ #pragma GCC diagnostic pop diff --git a/onnxruntime/contrib_ops/cpu/expand_dims.cc b/onnxruntime/contrib_ops/cpu/expand_dims.cc index 03b9d9418e..b3772c2a3b 100644 --- a/onnxruntime/contrib_ops/cpu/expand_dims.cc +++ b/onnxruntime/contrib_ops/cpu/expand_dims.cc @@ -2,8 +2,6 @@ // Licensed under the MIT License. #include "expand_dims.h" -#include "onnx/defs/schema.h" -#include "sample.h" namespace onnxruntime { namespace contrib { @@ -16,5 +14,5 @@ ONNX_CPU_OPERATOR_TYPED_MS_KERNEL( .TypeConstraint("T", DataTypeImpl::AllTensorTypes()) .TypeConstraint("axis", DataTypeImpl::GetTensorType()), contrib::ExpandDims); -} // namespace contrib +} // namespace contrib } // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cpu/sample.cc b/onnxruntime/contrib_ops/cpu/sample.cc index cd359d2e86..37b39a9d45 100644 --- a/onnxruntime/contrib_ops/cpu/sample.cc +++ b/onnxruntime/contrib_ops/cpu/sample.cc @@ -2,7 +2,6 @@ // Licensed under the MIT License. #include "sample.h" -#include "onnx/defs/schema.h" namespace onnxruntime { namespace contrib { @@ -13,5 +12,5 @@ ONNX_CPU_OPERATOR_TYPED_MS_KERNEL( float, KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()).MayInplace(0, 0), contrib::SampleOp); -} // namespace contrib +} // namespace contrib } // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cpu/tokenizer.cc b/onnxruntime/contrib_ops/cpu/tokenizer.cc index 8ba6a0a9af..8613233738 100644 --- a/onnxruntime/contrib_ops/cpu/tokenizer.cc +++ b/onnxruntime/contrib_ops/cpu/tokenizer.cc @@ -2,12 +2,9 @@ // Licensed under the MIT License. #include "core/common/common.h" +#include "core/common/utf8_util.h" #include "core/framework/tensor.h" #include "core/framework/op_kernel.h" -#include "core/graph/onnx_protobuf.h" -#include "onnx/defs/schema.h" - -#include "core/common/utf8_util.h" #include "re2/re2.h" namespace onnxruntime { diff --git a/onnxruntime/core/framework/kernel_registry.cc b/onnxruntime/core/framework/kernel_registry.cc index 3774cbf4d8..da5096e430 100644 --- a/onnxruntime/core/framework/kernel_registry.cc +++ b/onnxruntime/core/framework/kernel_registry.cc @@ -9,6 +9,7 @@ using namespace ::onnxruntime::common; namespace onnxruntime { +#if !defined(ORT_MINIMAL_BUILD) namespace { // Traverses the node's formal parameters and calls TraverseFn with the formal // parameter and its associated TypeProto. @@ -199,6 +200,8 @@ bool KernelRegistry::VerifyKernelDef(const onnxruntime::Node& node, return true; } +#endif // !defined(ORT_MINIMAL_BUILD) + Status KernelRegistry::Register(KernelDefBuilder& kernel_builder, const KernelCreateFn& kernel_creator) { return Register(KernelCreateInfo(kernel_builder.Build(), kernel_creator)); @@ -266,15 +269,20 @@ Status KernelRegistry::TryFindKernel(const onnxruntime::Node& node, auto range = kernel_creator_fn_map_.equal_range(GetMapKey(node.OpType(), node.Domain(), expected_provider)); std::vector verify_kernel_def_error_strs; + for (auto i = range.first; i != range.second; ++i) { +#if !defined(ORT_MINIMAL_BUILD) std::string error_str; if (VerifyKernelDef(node, *i->second.kernel_def, error_str)) { *out = &i->second; return Status::OK(); } verify_kernel_def_error_strs.push_back(error_str); +#else +// TODO: Add hash based lookup +#endif } - *out = nullptr; + if (!verify_kernel_def_error_strs.empty()) { std::ostringstream oss; oss << "Op with name (" << node.Name() << ")" @@ -283,6 +291,8 @@ Status KernelRegistry::TryFindKernel(const onnxruntime::Node& node, << " Encountered following errors: (" << ToString(verify_kernel_def_error_strs) << ")"; return Status(ONNXRUNTIME, FAIL, oss.str()); } + + *out = nullptr; return Status(ONNXRUNTIME, FAIL, "Kernel not found"); } diff --git a/onnxruntime/core/framework/kernel_registry_manager.cc b/onnxruntime/core/framework/kernel_registry_manager.cc index e7823f3722..737746cd09 100644 --- a/onnxruntime/core/framework/kernel_registry_manager.cc +++ b/onnxruntime/core/framework/kernel_registry_manager.cc @@ -3,10 +3,13 @@ #include "core/framework/kernel_registry_manager.h" #include "core/framework/kernel_registry.h" -#include "core/framework/customregistry.h" #include "core/framework/execution_providers.h" #include "core/framework/session_state.h" +#if !defined(ORT_MINIMAL_BUILD) +#include "core/framework/customregistry.h" +#endif + using namespace ONNX_NAMESPACE; using namespace ::onnxruntime::common; namespace onnxruntime { diff --git a/onnxruntime/core/framework/kernel_registry_manager.h b/onnxruntime/core/framework/kernel_registry_manager.h index d17dfc6f76..205ddae975 100644 --- a/onnxruntime/core/framework/kernel_registry_manager.h +++ b/onnxruntime/core/framework/kernel_registry_manager.h @@ -8,7 +8,6 @@ #include #include "core/common/status.h" #include "core/graph/graph_viewer.h" -#include "core/framework/customregistry.h" #include "core/platform/ort_mutex.h" namespace onnxruntime { diff --git a/onnxruntime/core/framework/op_node_proto_helper.cc b/onnxruntime/core/framework/op_node_proto_helper.cc index 8e39083593..b417a52466 100644 --- a/onnxruntime/core/framework/op_node_proto_helper.cc +++ b/onnxruntime/core/framework/op_node_proto_helper.cc @@ -1,16 +1,13 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include "core/common/exceptions.h" -#include "core/common/status.h" -#include "core/common/logging/logging.h" -#include "core/graph/graph_viewer.h" -#include "core/graph/op.h" -#include "onnx/defs/schema.h" +#include "core/common/common.h" #include "core/framework/op_kernel.h" #include "core/framework/tensorprotoutils.h" -#include "onnx/defs/schema.h" +#include "core/graph/onnx_protobuf.h" +#include "core/graph/op.h" #include "gsl/gsl" + using namespace ONNX_NAMESPACE; using namespace ::onnxruntime::common; namespace onnxruntime { @@ -49,7 +46,7 @@ inline bool HasTyped(const AttributeProto* attr) { if (!attr) { \ return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "No attribute with name:'", name, "'is defined."); \ } \ - if (!HasTyped(attr)) { \ + if (!HasTyped(attr)) { \ return Status(ONNXRUNTIME, FAIL, "Attibute name and type don't match"); \ } else { \ *value = static_cast(attr->type()); \ @@ -87,6 +84,7 @@ inline bool HasTyped(const AttributeProto* attr) { return Status::OK(); \ } +#if !defined(ORT_MINIMAL_BUILD) #define ORT_DEFINE_GET_ATTR_SPECIALIZATIONS(type, list) \ ORT_DEFINE_GET_ATTR(ProtoHelperNodeContext, type, list) \ ORT_DEFINE_GET_ATTR(InferenceContext, type, list) @@ -94,6 +92,13 @@ inline bool HasTyped(const AttributeProto* attr) { #define ORT_DEFINE_GET_ATTRS_SPECIALIZATIONS(type, list) \ ORT_DEFINE_GET_ATTRS(ProtoHelperNodeContext, type, list) \ ORT_DEFINE_GET_ATTRS(InferenceContext, type, list) +#else +#define ORT_DEFINE_GET_ATTR_SPECIALIZATIONS(type, list) \ + ORT_DEFINE_GET_ATTR(ProtoHelperNodeContext, type, list) + +#define ORT_DEFINE_GET_ATTRS_SPECIALIZATIONS(type, list) \ + ORT_DEFINE_GET_ATTRS(ProtoHelperNodeContext, type, list) +#endif ORT_DEFINE_GET_ATTR_SPECIALIZATIONS(float, f) ORT_DEFINE_GET_ATTR_SPECIALIZATIONS(int64_t, i) @@ -172,6 +177,8 @@ bool OpNodeProtoHelper::HasPrimitiveAttribute(AttributeProto_AttributeTy } template class OpNodeProtoHelper; +#if !defined(ORT_MINIMAL_BUILD) template class OpNodeProtoHelper; +#endif } // namespace onnxruntime diff --git a/onnxruntime/core/framework/provider_bridge_ort.cc b/onnxruntime/core/framework/provider_bridge_ort.cc index 539db10c3f..5ae24a2c1b 100644 --- a/onnxruntime/core/framework/provider_bridge_ort.cc +++ b/onnxruntime/core/framework/provider_bridge_ort.cc @@ -4,24 +4,27 @@ // This is the Onnxruntime side of the bridge to allow providers to be built as a DLL // It implements onnxruntime::ProviderHost -#include "core/framework/data_types.h" #include "core/framework/allocatormgr.h" +#include "core/framework/compute_capability.h" +#include "core/framework/data_types.h" +#include "core/framework/data_transfer_manager.h" +#include "core/framework/execution_provider.h" +#include "core/framework/kernel_registry.h" +#include "core/graph/model.h" +#include "core/platform/env.h" #include "core/providers/dnnl/dnnl_provider_factory.h" #include "core/session/inference_session.h" -#include "core/providers/tensorrt/tensorrt_provider_factory.h" +#include "core/session/abi_session_options_impl.h" +#include "core/session/ort_apis.h" + #ifdef USE_TENSORRT +#include "core/providers/tensorrt/tensorrt_provider_factory.h" #include "core/providers/cuda/cuda_allocator.h" #include "core/providers/cuda/gpu_data_transfer.h" #include "core/providers/cuda/math/unary_elementwise_ops_impl.h" #include "core/providers/cuda/cuda_common.h" #endif -#include "core/session/abi_session_options_impl.h" -#include "core/session/ort_apis.h" -#include "core/platform/env.h" -#include "core/graph/model.h" -#include "core/framework/data_transfer_manager.h" -#include "core/framework/compute_capability.h" -#include "core/framework/execution_provider.h" + #define PROVIDER_BRIDGE_ORT #include "core/providers/shared_library/provider_interfaces.h" #include "onnx/common/stl_backports.h" diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index e514e21aaf..36963944ab 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -7,8 +7,8 @@ #include "core/graph/contrib_ops/contrib_defs.h" #include "core/graph/contrib_ops/nchwc_schema_defs.h" #include "core/graph/contrib_ops/range_schema_defs.h" +#include "core/graph/onnx_protobuf.h" #include "core/graph/op.h" -#include "onnx/defs/schema.h" #include "onnx/defs/shape_inference.h" #include "onnx/defs/tensor_proto_util.h" #include "core/mlas/inc/mlas.h" diff --git a/onnxruntime/core/graph/featurizers_ops/featurizers_defs.cc b/onnxruntime/core/graph/featurizers_ops/featurizers_defs.cc index eac1d719b8..8648a4a3eb 100644 --- a/onnxruntime/core/graph/featurizers_ops/featurizers_defs.cc +++ b/onnxruntime/core/graph/featurizers_ops/featurizers_defs.cc @@ -3,11 +3,9 @@ #include "core/graph/constants.h" #include "core/graph/featurizers_ops/featurizers_defs.h" +#include "core/graph/onnx_protobuf.h" #include "core/graph/op.h" -#include "onnx/defs/schema.h" -#include "onnx/defs/shape_inference.h" - #define MS_FEATURIZERS_OPERATOR_SCHEMA(name) MS_FEATURIZERS_OPERATOR_SCHEMA_UNIQ_HELPER(__COUNTER__, name) #define MS_FEATURIZERS_OPERATOR_SCHEMA_UNIQ_HELPER(Counter, name) MS_FEATURIZERS_OPERATOR_SCHEMA_UNIQ(Counter, name) diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index 3b524240f1..82d4b7b432 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -21,12 +21,12 @@ #include "core/graph/graph_viewer.h" #include "core/graph/indexed_sub_graph.h" #include "core/graph/model.h" -#include "core/graph/schema_registry.h" #include "core/graph/op.h" #if !defined(ORT_MINIMAL_BUILD) #include "core/graph/function.h" #include "core/graph/function_impl.h" +#include "core/graph/schema_registry.h" #include "onnx/checker.h" using namespace ONNX_NAMESPACE::checker; #endif @@ -1580,7 +1580,7 @@ class InferenceContextImpl : public ONNX_NAMESPACE::InferenceContext { // only return data if it's for a constant initializer. checks for outer scope initializers // if this is a subgraph and the name isn't found locally. - const TensorProto* initializer = graph_utils::GetConstantInitializer(graph_, def->Name(), true); + const TensorProto* initializer = graph_.GetConstantInitializer(def->Name(), true); return initializer; } @@ -2392,6 +2392,31 @@ void Graph::CleanAllInitializedTensors() noexcept { } } +const ONNX_NAMESPACE::TensorProto* Graph::GetConstantInitializer(const std::string& initializer_name, + bool check_outer_scope) const { + const ONNX_NAMESPACE::TensorProto* initializer = nullptr; + if (GetInitializedTensor(initializer_name, initializer)) { + if (CanOverrideInitializer()) { + const auto& graph_inputs = GetInputsIncludingInitializers(); + bool is_constant = std::none_of(graph_inputs.cbegin(), graph_inputs.cend(), + [&initializer_name](const NodeArg* input) { + return input->Name() == initializer_name; + }); + + if (!is_constant) { + initializer = nullptr; + } + } + } else if (check_outer_scope && IsSubgraph()) { + // make sure there's not a local value with the same name. if there is it shadows any initializer in outer scope. + if (IsOuterScopeValue(initializer_name)) { + initializer = parent_graph_->GetConstantInitializer(initializer_name, check_outer_scope); + } + } + + return initializer; +} + #if !defined(ORT_MINIMAL_BUILD) void Graph::AddValueInfo(const NodeArg* new_value_info) { for (const auto* info : value_info_) { diff --git a/onnxruntime/core/graph/graph_utils.cc b/onnxruntime/core/graph/graph_utils.cc index 9100d3685f..a4370960fe 100644 --- a/onnxruntime/core/graph/graph_utils.cc +++ b/onnxruntime/core/graph/graph_utils.cc @@ -489,27 +489,7 @@ bool IsGraphInput(const Graph& graph, const NodeArg* input) { const ONNX_NAMESPACE::TensorProto* GetConstantInitializer(const Graph& graph, const std::string& initializer_name, bool check_outer_scope) { - const ONNX_NAMESPACE::TensorProto* initializer = nullptr; - if (graph.GetInitializedTensor(initializer_name, initializer)) { - if (graph.CanOverrideInitializer()) { - const auto& graph_inputs = graph.GetInputsIncludingInitializers(); - bool is_constant = std::none_of(graph_inputs.cbegin(), graph_inputs.cend(), - [&initializer_name](const NodeArg* input) { - return input->Name() == initializer_name; - }); - - if (!is_constant) { - initializer = nullptr; - } - } - } else if (check_outer_scope && graph.IsSubgraph()) { - // make sure there's not a local value with the same name. if there is it shadows any initializer in outer scope. - if (graph.IsOuterScopeValue(initializer_name)) { - initializer = GetConstantInitializer(*graph.ParentGraph(), initializer_name, check_outer_scope); - } - } - - return initializer; + return graph.GetConstantInitializer(initializer_name, check_outer_scope); } bool IsInitializer(const Graph& graph, const std::string& name, bool check_outer_scope) { diff --git a/onnxruntime/core/graph/graph_viewer.cc b/onnxruntime/core/graph/graph_viewer.cc index 100e629a28..aa29caef9b 100644 --- a/onnxruntime/core/graph/graph_viewer.cc +++ b/onnxruntime/core/graph/graph_viewer.cc @@ -8,8 +8,6 @@ #include "core/graph/graph_viewer.h" -#include "core/graph/graph_utils.h" - namespace onnxruntime { bool NodeCompare::operator()(const Node* n1, const Node* n2) const { @@ -115,7 +113,7 @@ bool GraphViewer::IsSubgraph() const { } bool GraphViewer::IsConstantInitializer(const std::string& name, bool check_outer_scope) const { - return graph_utils::IsConstantInitializer(*graph_, name, check_outer_scope); + return graph_->GetConstantInitializer(name, check_outer_scope) != nullptr; } } // namespace onnxruntime diff --git a/onnxruntime/core/graph/model.cc b/onnxruntime/core/graph/model.cc index 18aca63ab1..ca778c83d1 100644 --- a/onnxruntime/core/graph/model.cc +++ b/onnxruntime/core/graph/model.cc @@ -20,7 +20,10 @@ #include "gsl/gsl" #include "core/platform/env.h" + +#if !defined(ORT_MINIMAL_BUILD) #include "core/graph/schema_registry.h" +#endif using namespace ONNX_NAMESPACE; using namespace onnxruntime; diff --git a/onnxruntime/core/providers/cpu/nn/string_normalizer.cc b/onnxruntime/core/providers/cpu/nn/string_normalizer.cc index 12930a0a1f..92f87792ca 100644 --- a/onnxruntime/core/providers/cpu/nn/string_normalizer.cc +++ b/onnxruntime/core/providers/cpu/nn/string_normalizer.cc @@ -2,14 +2,13 @@ // Licensed under the MIT License. #include "string_normalizer.h" -#include "onnx/defs/schema.h" #include "core/common/common.h" #include "core/framework/tensor.h" #ifdef _MSC_VER #include #include -#elif defined (__APPLE__) or defined (__ANDROID__) +#elif defined(__APPLE__) or defined(__ANDROID__) #include #else #include @@ -77,7 +76,7 @@ using Utf8Converter = std::wstring_convert>; const std::string default_locale("en-US"); -#else // MS_VER +#else // MS_VER class Locale { public: @@ -112,7 +111,6 @@ using Utf8Converter = std::wstring_convert>; // All others (Linux) class Utf8Converter { public: - Utf8Converter(const std::string&, const std::wstring&) {} std::wstring from_bytes(const std::string& s) const { @@ -182,11 +180,11 @@ class Utf8Converter { } }; -#endif // __APPLE__ +#endif // __APPLE__ -const std::string default_locale("en_US.UTF-8"); // All non-MS +const std::string default_locale("en_US.UTF-8"); // All non-MS -#endif // MS_VER +#endif // MS_VER template Status CopyCaseAction(ForwardIter first, ForwardIter end, OpKernelContext* ctx, diff --git a/onnxruntime/core/providers/cpu/nn/tfidfvectorizer.cc b/onnxruntime/core/providers/cpu/nn/tfidfvectorizer.cc index 1bb0ce9d6c..5f552e7f05 100644 --- a/onnxruntime/core/providers/cpu/nn/tfidfvectorizer.cc +++ b/onnxruntime/core/providers/cpu/nn/tfidfvectorizer.cc @@ -2,7 +2,6 @@ // Licensed under the MIT License. #include "tfidfvectorizer.h" -#include "onnx/defs/schema.h" #include "core/common/common.h" #include "core/framework/tensor.h" #include "core/platform/threadpool.h" diff --git a/onnxruntime/core/providers/cpu/tensor/reverse_sequence.cc b/onnxruntime/core/providers/cpu/tensor/reverse_sequence.cc index d67fa7f324..8fbe0f939c 100644 --- a/onnxruntime/core/providers/cpu/tensor/reverse_sequence.cc +++ b/onnxruntime/core/providers/cpu/tensor/reverse_sequence.cc @@ -2,7 +2,6 @@ // Licensed under the MIT License. #include "reverse_sequence.h" -#include "onnx/defs/schema.h" // there's no way to use a raw pointer as the copy destination with std::copy_n // (which gsl::copy uses with span::data() which returns a raw pointer) with the 14.11 toolset diff --git a/onnxruntime/core/session/custom_ops.cc b/onnxruntime/core/session/custom_ops.cc index 6f29cc2b83..7389c79c18 100644 --- a/onnxruntime/core/session/custom_ops.cc +++ b/onnxruntime/core/session/custom_ops.cc @@ -7,7 +7,6 @@ #include "core/graph/onnx_protobuf.h" #include "core/session/inference_session.h" #include "core/session/ort_apis.h" -#include "core/framework/customregistry.h" #include "core/framework/data_types.h" #include "core/framework/op_kernel_info.h" #include "core/framework/op_kernel_context_internal.h" @@ -68,6 +67,9 @@ ORT_API_STATUS_IMPL(OrtApis::KernelInfoGetAttribute_string, _In_ const OrtKernel return onnxruntime::ToOrtStatus(status); } +#if !defined(ORT_MINIMAL_BUILD) +#include "core/framework/customregistry.h" + namespace onnxruntime { struct CustomOpKernel : OpKernel { @@ -149,3 +151,5 @@ common::Status CreateCustomRegistry(const std::vector& op_do } } // namespace onnxruntime + +#endif // !defined(ORT_MINIMAL_BUILD) diff --git a/onnxruntime/core/session/environment.cc b/onnxruntime/core/session/environment.cc index b8fbb162a3..96f797f129 100644 --- a/onnxruntime/core/session/environment.cc +++ b/onnxruntime/core/session/environment.cc @@ -5,9 +5,13 @@ #include "core/framework/allocatormgr.h" #include "core/graph/constants.h" #include "core/graph/op.h" +#if !defined(ORT_MINIMAL_BUILD) #include "onnx/defs/operator_sets.h" #include "onnx/defs/operator_sets_ml.h" +#if defined(ENABLE_TRAINING) #include "onnx/defs/operator_sets_training.h" +#endif +#endif #ifndef DISABLE_CONTRIB_OPS #include "core/graph/contrib_ops/contrib_defs.h" #endif diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index 4df0e21648..ed8cd6fc3a 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -12,43 +12,46 @@ #include #include "core/common/logging/logging.h" -#include "core/platform/threadpool.h" -#include "core/graph/graph_viewer.h" -#include "core/graph/graph_utils.h" -#include "core/graph/model.h" #include "core/framework/allocatormgr.h" -#include "core/framework/customregistry.h" -#include "core/session/environment.h" #include "core/framework/error_code_helper.h" #include "core/framework/execution_frame.h" #include "core/framework/feeds_fetches_manager.h" #include "core/framework/graph_partitioner.h" #include "core/framework/kernel_def_builder.h" #include "core/framework/kernel_registry.h" -#include "core/framework/ort_value_pattern_planner.h" #include "core/framework/mldata_type_utils.h" -#include "core/framework/op_kernel_context_internal.h" #include "core/framework/TensorSeq.h" #include "core/framework/tensorprotoutils.h" #include "core/framework/tensor_type_and_shape.h" +#include "core/framework/op_kernel_context_internal.h" +#include "core/framework/ort_value_pattern_planner.h" #include "core/framework/utils.h" +#include "core/graph/graph_viewer.h" +#include "core/graph/graph_utils.h" +#include "core/graph/model.h" #include "core/optimizer/transformer_memcpy.h" #include "core/optimizer/graph_transformer.h" #include "core/optimizer/insert_cast_transformer.h" +#include "core/optimizer/rule_based_graph_transformer.h" +#include "core/optimizer/graph_transformer_utils.h" +#include "core/platform/Barrier.h" +#include "core/platform/ort_mutex.h" +#include "core/platform/threadpool.h" #include "core/providers/cpu/controlflow/utils.h" #include "core/providers/cpu/cpu_execution_provider.h" #ifdef USE_DML // TODO: This is necessary for the workaround in TransformGraph #include "core/providers/dml/DmlExecutionProvider/src/GraphTransformer.h" #endif +#include "core/session/environment.h" #include "core/session/IOBinding.h" -#include "core/session/custom_ops.h" -#include "core/util/protobuf_parsing_utils.h" -#include "core/optimizer/rule_based_graph_transformer.h" -#include "core/optimizer/graph_transformer_utils.h" -#include "core/util/thread_utils.h" #include "core/session/inference_session_utils.h" -#include "core/platform/ort_mutex.h" -#include "core/platform/Barrier.h" +#include "core/util/protobuf_parsing_utils.h" +#include "core/util/thread_utils.h" + +#if !defined(ORT_MINIMAL_BUILD) +#include "core/framework/customregistry.h" +#include "core/session/custom_ops.h" +#endif using namespace ONNX_NAMESPACE; diff --git a/onnxruntime/python/onnxruntime_pybind_state.cc b/onnxruntime/python/onnxruntime_pybind_state.cc index e9bdcbea4d..84471ab68a 100644 --- a/onnxruntime/python/onnxruntime_pybind_state.cc +++ b/onnxruntime/python/onnxruntime_pybind_state.cc @@ -9,16 +9,17 @@ #define PY_ARRAY_UNIQUE_SYMBOL onnxruntime_python_ARRAY_API #include -#include "core/framework/data_transfer_utils.h" -#include "core/framework/data_types_internal.h" -#include "core/framework/tensorprotoutils.h" -#include "core/graph/graph_viewer.h" #include "core/common/logging/logging.h" #include "core/common/logging/severity.h" -#include "core/framework/TensorSeq.h" +#include "core/framework/bfc_arena.h" +#include "core/framework/data_transfer_utils.h" +#include "core/framework/data_types_internal.h" +#include "core/framework/kernel_registry.h" #include "core/framework/random_seed.h" #include "core/framework/session_options.h" -#include "core/framework/bfc_arena.h" +#include "core/framework/tensorprotoutils.h" +#include "core/framework/TensorSeq.h" +#include "core/graph/graph_viewer.h" #include "core/session/IOBinding.h" #if USE_CUDA diff --git a/onnxruntime/test/framework/allocation_planner_test.cc b/onnxruntime/test/framework/allocation_planner_test.cc index 96535ddc3c..41bcff9df2 100644 --- a/onnxruntime/test/framework/allocation_planner_test.cc +++ b/onnxruntime/test/framework/allocation_planner_test.cc @@ -5,14 +5,17 @@ #include #include #include -#include "core/framework/session_state.h" -#include "core/graph/model.h" #include "gtest/gtest.h" + +#include "core/framework/session_state.h" +#include "core/framework/kernel_registry.h" #include "core/framework/op_kernel.h" #include "test/framework/model_builder_utils.h" #include "core/framework/allocation_planner.h" -#include "core/util/thread_utils.h" +#include "core/graph/model.h" #include "core/providers/cpu/cpu_execution_provider.h" +#include "core/util/thread_utils.h" + #include "test/test_environment.h" #include "test/util/include/asserts.h" diff --git a/onnxruntime/test/framework/allocator_test.cc b/onnxruntime/test/framework/allocator_test.cc index 1f4600f298..061d8b802f 100644 --- a/onnxruntime/test/framework/allocator_test.cc +++ b/onnxruntime/test/framework/allocator_test.cc @@ -3,6 +3,7 @@ #include "core/framework/allocatormgr.h" #include "core/framework/allocator.h" + #include "test_utils.h" #include "gtest/gtest.h" diff --git a/onnxruntime/test/framework/session_state_test.cc b/onnxruntime/test/framework/session_state_test.cc index fecb8f194c..a11e8fb232 100644 --- a/onnxruntime/test/framework/session_state_test.cc +++ b/onnxruntime/test/framework/session_state_test.cc @@ -5,6 +5,7 @@ #include "core/framework/execution_providers.h" #include "core/framework/graph_partitioner.h" +#include "core/framework/kernel_registry.h" #include "core/framework/op_kernel.h" #include "core/framework/session_state.h" #include "core/graph/graph_utils.h" @@ -259,7 +260,7 @@ TEST_P(SessionStatePrepackingTest, PrePackingTest) { SessionOptions sess_options; bool use_prepacking = GetParam(); sess_options.session_configurations[ORT_SESSION_OPTIONS_CONFIG_DISABLEPREPACKING] = use_prepacking ? "0" : "1"; - ASSERT_STATUS_OK(session_state.FinalizeSessionState(std::basic_string(), + ASSERT_STATUS_OK(session_state.FinalizeSessionState(std::basic_string(), kernel_registry_manager, sess_options)); From 728e886bba0e5dbf264bcf921b2ec86b19fea981 Mon Sep 17 00:00:00 2001 From: Scott McKay Date: Sun, 23 Aug 2020 14:39:07 +1000 Subject: [PATCH 3/4] Add kernel def hash logic for minimal build (#4891) * Add hash based lookup of kernels --- include/onnxruntime/core/common/common.h | 7 + .../core/framework/kernel_def_builder.h | 32 ++++ .../core/framework/kernel_registry.h | 27 ++- .../core/framework/execution_provider.cc | 4 + onnxruntime/core/framework/kernel_registry.cc | 161 ++++++++++-------- .../core/framework/kernel_registry_manager.cc | 17 +- .../core/framework/kernel_registry_manager.h | 16 +- 7 files changed, 182 insertions(+), 82 deletions(-) diff --git a/include/onnxruntime/core/common/common.h b/include/onnxruntime/core/common/common.h index f6ff7d3fdf..eda0f429c6 100644 --- a/include/onnxruntime/core/common/common.h +++ b/include/onnxruntime/core/common/common.h @@ -242,4 +242,11 @@ inline std::wstring ToWideString(const std::wstring& s) { return s; } inline std::string ToWideString(const std::string& s) { return s; } #endif +// from http://www.open-std.org/jtc1/sc22/wg21/docs/papers/2014/n3876.pdf +template +inline void HashCombine(std::uint64_t& seed, const T& v) { + std::hash hasher; + seed ^= hasher(v) + 0x9e3779b9 + (seed << 6) + (seed >> 2); +} + } // namespace onnxruntime diff --git a/include/onnxruntime/core/framework/kernel_def_builder.h b/include/onnxruntime/core/framework/kernel_def_builder.h index 1b7027668b..0b47e6743c 100644 --- a/include/onnxruntime/core/framework/kernel_def_builder.h +++ b/include/onnxruntime/core/framework/kernel_def_builder.h @@ -88,9 +88,37 @@ class KernelDef { bool IsConflict(const KernelDef& other) const; + uint64_t GetHash() const noexcept { + // if we need to support different hash versions we can update CalculateHash to take a version number + // and calculate any non-default versions dynamically. we only use this during kernel lookup so + // it's not performance critical + return hash_; + } + private: friend class KernelDefBuilder; + // call once the KernelDef has been built + void CalculateHash() { + // use name, start/end, domain, provider and the type constraints. + // we wouldn't have two kernels that only differed by the inplace or alias info or memory types. + // currently nothing sets exec_queue_id either (and would assumably be a runtime thing and not part of the base + // kernel definition) + hash_ = 0; // reset in case this is called multiple times + HashCombine(hash_, op_name_); + HashCombine(hash_, op_since_version_start_); + HashCombine(hash_, op_since_version_end_); + HashCombine(hash_, op_domain_); + HashCombine(hash_, provider_type_); + for (const auto& key_value : type_constraints_) { + HashCombine(hash_, key_value.first); + for (const auto& data_type : key_value.second) { + // need to construct a std::string so it doesn't hash the address of a const char* + HashCombine(hash_, std::string(DataTypeImpl::ToString(data_type))); + } + } + } + // The operator name supported by <*this> kernel.. std::string op_name_; @@ -128,6 +156,9 @@ class KernelDef { OrtMemType default_inputs_mem_type_{OrtMemTypeDefault}; // Default memory type for all outputs OrtMemType default_outputs_mem_type_{OrtMemTypeDefault}; + + // hash of kernel definition for lookup in minimal build + uint64_t hash_ = 0; }; class KernelDefBuilder { @@ -283,6 +314,7 @@ class KernelDefBuilder { Return the kernel definition, passing ownership of the KernelDef to the caller */ std::unique_ptr Build() { + kernel_def_->CalculateHash(); return std::move(kernel_def_); } diff --git a/include/onnxruntime/core/framework/kernel_registry.h b/include/onnxruntime/core/framework/kernel_registry.h index ca92da91c6..0542183d6d 100644 --- a/include/onnxruntime/core/framework/kernel_registry.h +++ b/include/onnxruntime/core/framework/kernel_registry.h @@ -19,6 +19,14 @@ class KernelRegistry { Status Register(KernelCreateInfo&& create_info) ORT_MUST_USE_RESULT; +#if !defined(ORT_MINIMAL_BUILD) + static bool HasImplementationOf(const KernelRegistry& r, const onnxruntime::Node& node, + onnxruntime::ProviderType exec_provider) { + const KernelCreateInfo* info; + Status st = r.TryFindKernel(node, exec_provider, &info); + return st.IsOK(); + } + // factory functions should always return a unique_ptr for maximum flexibility // for its clients unless the factory is managing the lifecycle of the pointer // itself. @@ -30,15 +38,16 @@ class KernelRegistry { std::unique_ptr& op_kernel) const ORT_MUST_USE_RESULT; // Check if an execution provider can create kernel for a node and return the kernel if so - Status TryFindKernel(const onnxruntime::Node& node, - onnxruntime::ProviderType exec_provider, const KernelCreateInfo** out) const; + Status TryFindKernel(const onnxruntime::Node& node, onnxruntime::ProviderType exec_provider, + const KernelCreateInfo** out) const; - static bool HasImplementationOf(const KernelRegistry& r, const onnxruntime::Node& node, - onnxruntime::ProviderType exec_provider) { - const KernelCreateInfo* info; - Status st = r.TryFindKernel(node, exec_provider, &info); - return st.IsOK(); - } +#endif + + // Check if an execution provider can create kernel for a node and return the kernel if so. + // Kernel matching is via kernel_def_hash. + Status TryFindKernel(const onnxruntime::Node& node, onnxruntime::ProviderType exec_provider, + uint64_t kernel_def_hash, + const KernelCreateInfo** out) const; bool IsEmpty() const { return kernel_creator_fn_map_.empty(); } @@ -50,6 +59,7 @@ class KernelRegistry { #endif private: +#if !defined(ORT_MINIMAL_BUILD) // Check whether the types of inputs/outputs of the given node match the extra // type-constraints of the given kernel. This serves two purposes: first, to // select the right kernel implementation based on the types of the arguments @@ -67,6 +77,7 @@ class KernelRegistry { static bool VerifyKernelDef(const onnxruntime::Node& node, const KernelDef& kernel_def, std::string& error_str); +#endif static std::string GetMapKey(const std::string& op_name, const std::string& domain, const std::string& provider) { std::string key(op_name); diff --git a/onnxruntime/core/framework/execution_provider.cc b/onnxruntime/core/framework/execution_provider.cc index 4c34291315..f771f33f00 100644 --- a/onnxruntime/core/framework/execution_provider.cc +++ b/onnxruntime/core/framework/execution_provider.cc @@ -29,6 +29,7 @@ std::vector> IExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph, const std::vector& kernel_registries) const { std::vector> result; +#if !defined(ORT_MINIMAL_BUILD) for (auto& node : graph.Nodes()) { for (auto registry : kernel_registries) { if (KernelRegistry::HasImplementationOf(*registry, node, Type())) { @@ -39,6 +40,9 @@ IExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph, } } } +#else + ORT_NOT_IMPLEMENTED("IExecutionProvider::GetCapability is not supported in this build."); +#endif return result; } diff --git a/onnxruntime/core/framework/kernel_registry.cc b/onnxruntime/core/framework/kernel_registry.cc index da5096e430..97dc6cbd47 100644 --- a/onnxruntime/core/framework/kernel_registry.cc +++ b/onnxruntime/core/framework/kernel_registry.cc @@ -200,8 +200,102 @@ bool KernelRegistry::VerifyKernelDef(const onnxruntime::Node& node, return true; } +Status KernelRegistry::TryCreateKernel(const onnxruntime::Node& node, + const IExecutionProvider& execution_provider, + const std::unordered_map& constant_initialized_tensors, + const OrtValueNameIdxMap& ort_value_name_idx_map, + const FuncManager& funcs_mgr, + const DataTransferManager& data_transfer_mgr, + /*out*/ std::unique_ptr& op_kernel) const { + const KernelCreateInfo* kernel_create_info = nullptr; + ORT_RETURN_IF_ERROR(TryFindKernel(node, execution_provider.Type(), &kernel_create_info)); + OpKernelInfo kernel_info(node, + *kernel_create_info->kernel_def, + execution_provider, + constant_initialized_tensors, + ort_value_name_idx_map, + funcs_mgr, + data_transfer_mgr); + op_kernel.reset(kernel_create_info->kernel_create_func(kernel_info)); + return Status::OK(); +} + +static std::string ToString(const std::vector& error_strs) { + std::ostringstream ostr; + std::for_each(std::begin(error_strs), std::end(error_strs), + [&ostr](const std::string& str) { ostr << str << "\n"; }); + return ostr.str(); +} + +Status KernelRegistry::TryFindKernel(const onnxruntime::Node& node, + onnxruntime::ProviderType exec_provider, + const KernelCreateInfo** out) const { + return TryFindKernel(node, exec_provider, uint64_t(0), out); +} #endif // !defined(ORT_MINIMAL_BUILD) +// It's often this function returns a failed status, but it is totally expected. +// It just means this registry doesn't have such a kernel, please search it elsewhere. +// if this function is called before graph partition, then node.provider is not set. +// In this case, the kernel's provider must equal to exec_provider +// otherwise, kernel_def.provider must equal to node.provider. exec_provider is ignored. +Status KernelRegistry::TryFindKernel(const onnxruntime::Node& node, + onnxruntime::ProviderType exec_provider, + uint64_t kernel_def_hash, + const KernelCreateInfo** out) const { + const auto& node_provider = node.GetExecutionProviderType(); + const auto& expected_provider = (node_provider.empty() ? exec_provider : node_provider); + + auto range = kernel_creator_fn_map_.equal_range(GetMapKey(node.OpType(), node.Domain(), expected_provider)); + *out = nullptr; + + // if we have a hash (ORT format model) use only that. + if (kernel_def_hash != 0) { + for (auto i = range.first; i != range.second; ++i) { + if (i->second.kernel_def->GetHash() == kernel_def_hash) { + *out = &i->second; + return Status::OK(); + } + } + + std::ostringstream oss; + oss << "Op with name (" << node.Name() << ")" + << " and type (" << node.OpType() << ")" + << " kernel not found in " << expected_provider << "." + << " No matching hash for " << kernel_def_hash; + + return Status(ONNXRUNTIME, FAIL, oss.str()); + } +#if !defined(ORT_MINIMAL_BUILD) + else { + std::vector verify_kernel_def_error_strs; + + for (auto i = range.first; i != range.second; ++i) { + std::string error_str; + if (VerifyKernelDef(node, *i->second.kernel_def, error_str)) { + *out = &i->second; + return Status::OK(); + } + verify_kernel_def_error_strs.push_back(error_str); + } + + if (!verify_kernel_def_error_strs.empty()) { + std::ostringstream oss; + oss << "Op with name (" << node.Name() << ")" + << " and type (" << node.OpType() << ")" + << " kernel is not supported in " << expected_provider << "." + << " Encountered following errors: (" << ToString(verify_kernel_def_error_strs) << ")"; + + return Status(ONNXRUNTIME, FAIL, oss.str()); + } + } + + return Status(ONNXRUNTIME, FAIL, "Kernel not found"); +#else + ORT_THROW("Kernel hash must be provided in minimal build.") +#endif +} + Status KernelRegistry::Register(KernelDefBuilder& kernel_builder, const KernelCreateFn& kernel_creator) { return Register(KernelCreateInfo(kernel_builder.Build(), kernel_creator)); @@ -229,71 +323,4 @@ Status KernelRegistry::Register(KernelCreateInfo&& create_info) { return Status::OK(); } -Status KernelRegistry::TryCreateKernel(const onnxruntime::Node& node, - const IExecutionProvider& execution_provider, - const std::unordered_map& constant_initialized_tensors, - const OrtValueNameIdxMap& ort_value_name_idx_map, - const FuncManager& funcs_mgr, - const DataTransferManager& data_transfer_mgr, - /*out*/ std::unique_ptr& op_kernel) const { - const KernelCreateInfo* kernel_create_info = nullptr; - ORT_RETURN_IF_ERROR(TryFindKernel(node, execution_provider.Type(), &kernel_create_info)); - OpKernelInfo kernel_info(node, - *kernel_create_info->kernel_def, - execution_provider, - constant_initialized_tensors, - ort_value_name_idx_map, - funcs_mgr, - data_transfer_mgr); - op_kernel.reset(kernel_create_info->kernel_create_func(kernel_info)); - return Status::OK(); -} - -static std::string ToString(const std::vector& error_strs) { - std::ostringstream ostr; - std::for_each(std::begin(error_strs), std::end(error_strs), - [&ostr](const std::string& str) { ostr << str << "\n"; }); - return ostr.str(); -} - -// It's often this function returns a failed status, but it is totally expected. -// It just means this registry doesn't have such a kernel, please search it elsewhere. -// if this function is called before graph partition, then node.provider is not set. -// In this case, the kernel's provider must equal to exec_provider -// otherwise, kernel_def.provider must equal to node.provider. exec_provider is ignored. - -Status KernelRegistry::TryFindKernel(const onnxruntime::Node& node, - onnxruntime::ProviderType exec_provider, const KernelCreateInfo** out) const { - const auto& node_provider = node.GetExecutionProviderType(); - const auto& expected_provider = (node_provider.empty() ? exec_provider : node_provider); - - auto range = kernel_creator_fn_map_.equal_range(GetMapKey(node.OpType(), node.Domain(), expected_provider)); - std::vector verify_kernel_def_error_strs; - - for (auto i = range.first; i != range.second; ++i) { -#if !defined(ORT_MINIMAL_BUILD) - std::string error_str; - if (VerifyKernelDef(node, *i->second.kernel_def, error_str)) { - *out = &i->second; - return Status::OK(); - } - verify_kernel_def_error_strs.push_back(error_str); -#else -// TODO: Add hash based lookup -#endif - } - - if (!verify_kernel_def_error_strs.empty()) { - std::ostringstream oss; - oss << "Op with name (" << node.Name() << ")" - << " and type (" << node.OpType() << ")" - << " kernel is not supported in " << expected_provider << "." - << " Encountered following errors: (" << ToString(verify_kernel_def_error_strs) << ")"; - return Status(ONNXRUNTIME, FAIL, oss.str()); - } - - *out = nullptr; - return Status(ONNXRUNTIME, FAIL, "Kernel not found"); -} - } // namespace onnxruntime diff --git a/onnxruntime/core/framework/kernel_registry_manager.cc b/onnxruntime/core/framework/kernel_registry_manager.cc index 737746cd09..b7e772701f 100644 --- a/onnxruntime/core/framework/kernel_registry_manager.cc +++ b/onnxruntime/core/framework/kernel_registry_manager.cc @@ -45,6 +45,7 @@ Status KernelRegistryManager::RegisterKernels(const ExecutionProviders& executio return Status::OK(); } +#if !defined(ORT_MINIMAL_BUILD) void KernelRegistryManager::RegisterKernelRegistry(std::shared_ptr kernel_registry) { if (nullptr == kernel_registry) { return; @@ -61,6 +62,12 @@ bool KernelRegistryManager::HasImplementationOf(const KernelRegistryManager& r, Status KernelRegistryManager::SearchKernelRegistry(const onnxruntime::Node& node, /*out*/ const KernelCreateInfo** kernel_create_info) const { + return SearchKernelRegistry(node, uint64_t(0), kernel_create_info); +} +#endif + +Status KernelRegistryManager::SearchKernelRegistry(const onnxruntime::Node& node, uint64_t kernel_def_hash, + /*out*/ const KernelCreateInfo** kernel_create_info) const { Status status; auto create_error_message = [&node, &status](const std::string& prefix) { @@ -77,10 +84,12 @@ Status KernelRegistryManager::SearchKernelRegistry(const onnxruntime::Node& node return Status(ONNXRUNTIME, FAIL, create_error_message("The node is not placed on any Execution Provider. ")); } +#if !defined(ORT_MINIMAL_BUILD) for (auto& registry : custom_kernel_registries_) { - status = registry->TryFindKernel(node, std::string(), kernel_create_info); + status = registry->TryFindKernel(node, std::string(), kernel_def_hash, kernel_create_info); if (status.IsOK()) return status; } +#endif KernelRegistry* p = nullptr; auto iter = provider_type_to_registry_.find(ptype); @@ -89,8 +98,10 @@ Status KernelRegistryManager::SearchKernelRegistry(const onnxruntime::Node& node } if (p != nullptr) { - status = p->TryFindKernel(node, std::string(), kernel_create_info); - if (status.IsOK()) return status; + status = p->TryFindKernel(node, std::string(), kernel_def_hash, kernel_create_info); + if (status.IsOK()) { + return status; + } } return Status(ONNXRUNTIME, NOT_IMPLEMENTED, create_error_message("Failed to find kernel for ")); diff --git a/onnxruntime/core/framework/kernel_registry_manager.h b/onnxruntime/core/framework/kernel_registry_manager.h index 205ddae975..0e34e1e57d 100644 --- a/onnxruntime/core/framework/kernel_registry_manager.h +++ b/onnxruntime/core/framework/kernel_registry_manager.h @@ -32,6 +32,8 @@ class KernelRegistryManager { // Register kernels from providers Status RegisterKernels(const ExecutionProviders& execution_providers) ORT_MUST_USE_RESULT; +#if !defined(ORT_MINIMAL_BUILD) + // The registry passed in this function has highest priority than anything already in this KernelRegistryManager, // and anything registered from RegisterKernels // For example, if you do: @@ -46,10 +48,6 @@ class KernelRegistryManager { Status SearchKernelRegistry(const onnxruntime::Node& node, /*out*/ const KernelCreateInfo** kernel_create_info) const; - std::unique_ptr CreateKernel(const onnxruntime::Node& node, const IExecutionProvider& execution_provider, - const SessionState& session_state, - const KernelCreateInfo& kernel_create_info) const ORT_MUST_USE_RESULT; - /** * Whether this node can be run on this provider */ @@ -70,6 +68,16 @@ class KernelRegistryManager { if (iter != provider_type_to_registry_.end()) result.push_back(iter->second.get()); return result; } +#endif + + Status SearchKernelRegistry(const onnxruntime::Node& node, + uint64_t kernel_def_hash, + /*out*/ const KernelCreateInfo** kernel_create_info) const; + + std::unique_ptr CreateKernel(const onnxruntime::Node& node, + const IExecutionProvider& execution_provider, + const SessionState& session_state, + const KernelCreateInfo& kernel_create_info) const ORT_MUST_USE_RESULT; ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(KernelRegistryManager); From eb05db5a2ab989a7696c51beeb566e2f13a8b85e Mon Sep 17 00:00:00 2001 From: Rayan-Krishnan Date: Sat, 22 Aug 2020 22:04:17 -0700 Subject: [PATCH 4/4] Fix OptimizerConfig params groups (#4877) * Copy samples to build folder and load models from there. Fix CI * This PR also includes a fix to path validation for save_as_onnx API * Add torchtext to CI for GPU training * Remove new frontend tests from CI Co-authored-by: Thiago Crepaldi --- .../python/experimental/optim/config.py | 9 +- .../python/experimental/orttrainer.py | 19 +- ...ttraining_test_orttrainer_bert_toy_onnx.py | 162 ++++++++++++++---- .../orttraining_test_orttrainer_frontend.py | 5 +- tools/ci_build/build.py | 2 - .../linux/docker/scripts/install_deps.sh | 2 +- 6 files changed, 156 insertions(+), 43 deletions(-) diff --git a/orttraining/orttraining/python/experimental/optim/config.py b/orttraining/orttraining/python/experimental/optim/config.py index bcc008aa76..62349cd561 100644 --- a/orttraining/orttraining/python/experimental/optim/config.py +++ b/orttraining/orttraining/python/experimental/optim/config.py @@ -53,7 +53,7 @@ class _OptimizerConfig(object): " and additional entries for custom hyper parameter values") for k, _ in group.items(): if k != 'params': - assert k in defaults, f"'params' has 'k' hyper parameter not present at 'defaults'" + assert k in defaults or k.replace("_coef", "") in defaults, f"'params' has {k} hyper parameter not present at 'defaults'" self.name = name self.lr = float(defaults['lr']) @@ -75,6 +75,9 @@ class _OptimizerConfig(object): if name not in param_group: param_group.setdefault(name, value) + if "lambda_coef" in param_group: + param_group["lambda"] = param_group.pop("lambda_coef") + self.params.append(param_group) @@ -160,7 +163,7 @@ class AdamConfig(_OptimizerConfig): defaults = {'lr': lr, 'alpha': alpha, 'beta': beta, - 'lambda_coef': lambda_coef, + 'lambda': lambda_coef, 'epsilon': epsilon, 'do_bias_correction': do_bias_correction, 'weight_decay_mode': weight_decay_mode} @@ -223,7 +226,7 @@ class LambConfig(_OptimizerConfig): defaults = {'lr': lr, 'alpha': alpha, 'beta': beta, - 'lambda_coef': lambda_coef, + 'lambda': lambda_coef, 'ratio_min': ratio_min, 'ratio_max': ratio_max, 'epsilon': epsilon, diff --git a/orttraining/orttraining/python/experimental/orttrainer.py b/orttraining/orttraining/python/experimental/orttrainer.py index 3e42a26bfe..fb2c5b1d41 100644 --- a/orttraining/orttraining/python/experimental/orttrainer.py +++ b/orttraining/orttraining/python/experimental/orttrainer.py @@ -525,12 +525,27 @@ class ORTTrainer(object): trainable_params.add(initializer.name) optimizer_attributes_map[initializer.name] = {} optimizer_int_attributes_map[initializer.name] = {} + not_in_param_groups = True for param_group in self.optim_config.params: if initializer.name not in param_group['params']: continue # keep looking for a matching param_group + not_in_param_groups = False for k, v in param_group.items(): - if k == 'params': - continue # 'params' is not a hyper parameter, skip it + # 'params' is not a hyper parameter, skip it. 'lr' per weight is not supported + if k == 'params' or k == 'lr': + continue + if isinstance(v, float): + optimizer_attributes_map[initializer.name][k] = v + elif isinstance(v, int): + optimizer_int_attributes_map[initializer.name][k] = v + else: + raise ValueError("Optimizer attributes must be either float or int.") + + # set default values for params not found in groups + if not_in_param_groups: + for k, v in self.optim_config.defaults.items(): + if k == 'lr': + continue if isinstance(v, float): optimizer_attributes_map[initializer.name][k] = v elif isinstance(v, int): diff --git a/orttraining/orttraining/test/python/orttraining_test_orttrainer_bert_toy_onnx.py b/orttraining/orttraining/test/python/orttraining_test_orttrainer_bert_toy_onnx.py index 9bf463e830..d41be439a0 100644 --- a/orttraining/orttraining/test/python/orttraining_test_orttrainer_bert_toy_onnx.py +++ b/orttraining/orttraining/test/python/orttraining_test_orttrainer_bert_toy_onnx.py @@ -70,6 +70,22 @@ def bert_model_description(dynamic_shape=True): return model_desc +def optimizer_parameters_mutiple_groups(model): + '''A method to assign different hyper parameters for different model parameter groups''' + + no_decay_keys = ["bias", "gamma", "beta", "LayerNorm"] + no_decay_param_group = [] + decay_param_group = [] + for initializer in model.graph.initializer: + if any(key in initializer.name for key in no_decay_keys): + no_decay_param_group.append(initializer.name) + else: + decay_param_group.append(initializer.name) + params = [{'params': no_decay_param_group, "alpha": 0.9, "beta": 0.999, "lambda_coef": 0.0, "epsilon": 1e-6}, + {'params': decay_param_group, "alpha": 0.9, "beta": 0.999, "lambda_coef": 0.01, "epsilon": 1e-6}] + return params + + def optimizer_parameters(model): '''A method to assign different hyper parameters for different model parameter groups''' @@ -78,7 +94,7 @@ def optimizer_parameters(model): for initializer in model.graph.initializer: if any(key in initializer.name for key in no_decay_keys): no_decay_param_group.append(initializer.name) - params = [{'params': no_decay_param_group, "alpha": 0.9, "beta": 0.999, "lambda_coef": 0.0, "epsilon": 1e-6}] + params = [{'params': no_decay_param_group, "alpha": 0.9, "beta": 0.999, "lambda_coef": 0.0, "epsilon": 1e-6, "do_bias_correction":False}] return params @@ -122,14 +138,6 @@ def legacy_model_params(lr, device = torch.device("cuda", 0)): learning_rate = torch.tensor([lr]).to(device) return (legacy_model_desc, learning_rate_description, learning_rate) - no_decay_keys = ["bias", "gamma", "beta", "LayerNorm"] - no_decay = any(no_decay_key in name for no_decay_key in no_decay_keys) - if no_decay: - return {"alpha": 0.9, "beta": 0.999, "lambda": 0.0, "epsilon": 1e-6} - else: - return {"alpha": 0.9, "beta": 0.999, "lambda": 0.01, "epsilon": 1e-6} - - def legacy_ort_trainer_learning_rate_description(): return Legacy_IODescription('Learning_Rate', [1, ], torch.float32) @@ -147,6 +155,23 @@ def legacy_bert_model_description(): next_sentence_labels_desc], [loss_desc]) +def legacy_optim_params_a(name): + return {"alpha": 0.9, "beta": 0.999, "lambda": 0.01, "epsilon": 1e-6, "do_bias_correction": False} + + +def legacy_optim_params_b(name): + params = ['bert.embeddings.LayerNorm.bias', 'bert.embeddings.LayerNorm.weight'] + if name in params: + return {"alpha": 0.9, "beta": 0.999, "lambda": 0.0, "epsilon": 1e-6, "do_bias_correction": False} + return {"alpha": 0.9, "beta": 0.999, "lambda": 0.01, "epsilon": 1e-6, "do_bias_correction": False} + + +def legacy_optim_params_c(name): + params_group = optimizer_parameters(load_bert_onnx_model()) + if name in params_group[0]['params']: + return {"alpha": 0.9, "beta": 0.999, "lambda": 0.0, "epsilon": 1e-6, "do_bias_correction": False} + return {"alpha": 0.9, "beta": 0.999, "lambda": 0.01, "epsilon": 1e-6, "do_bias_correction": False} + ############################################################################### # Testing starts here ######################################################### ############################################################################### @@ -209,34 +234,51 @@ def testToyBERTDeterministicCheck(expected_losses): @pytest.mark.parametrize("initial_lr, lr_scheduler, expected_learning_rates, expected_losses", [ (1.0, optim.lr_scheduler.ConstantWarmupLRScheduler,\ - [0.0, 0.2, 0.4, 0.6, 0.8, 1.0, 1.0, 1.0, 1.0, 1.0], - [10.988012313842773, 10.99213981628418, 18.11327362060547, 35.164119720458984, 40.68449401855469,\ - 24.585613250732422, 90.66416931152344, 97.23558044433594, 123.73894500732422, 185.28268432617188]), + [0.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0], + [10.988012313842773, 10.99213981628418, 120.79301452636719, 36.11647033691406, 95.83200073242188,\ + 221.2766571044922, 208.40316772460938, 279.5332946777344, 402.46380615234375, 325.79254150390625]), (0.5, optim.lr_scheduler.ConstantWarmupLRScheduler,\ - [0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.5, 0.5, 0.5, 0.5], - [10.988012313842773, 10.99213981628418, 14.942025184631348, 32.938804626464844, 15.188745498657227,\ - 31.730798721313477, 38.70293045043945, 29.556488037109375, 34.914039611816406, 43.36301040649414]), + [0.0, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5], + [10.988012313842773, 10.99213981628418, 52.69743347167969, 19.741533279418945, 83.88340759277344,\ + 126.39848327636719, 91.53898620605469, 63.62016296386719, 102.21206665039062, 180.1424560546875]), (1.0, optim.lr_scheduler.CosineWarmupLRScheduler,\ - [0.0, 0.2, 0.4, 0.6, 0.8, 1.0, 0.9045084971874737, 0.6545084971874737, 0.34549150281252633, 0.09549150281252633], - [10.988012313842773, 10.99213981628418, 18.11327362060547, 35.164119720458984, 40.68449401855469,\ - 24.585613250732422, 90.66416931152344, 95.54305267333984, 77.43124389648438, 111.71074676513672]), + [0.0, 0.9931806517013612, 0.9397368756032445, 0.8386407858128706, 0.7008477123264848, 0.5412896727361662,\ + 0.37725725642960045, 0.22652592093878665, 0.10542974530180327, 0.02709137914968268], + [10.988012313842773, 10.99213981628418, 120.6441650390625, 32.152557373046875, 89.63705444335938,\ + 138.8782196044922, 117.57748413085938, 148.01927185058594, 229.60403442382812, 110.2930908203125]), (1.0, optim.lr_scheduler.LinearWarmupLRScheduler,\ - [0.0, 0.2, 0.4, 0.6, 0.8, 1.0, 0.8, 0.6, 0.4, 0.2],\ - [10.988012313842773, 10.99213981628418, 18.11327362060547, 35.164119720458984, 40.68449401855469, 24.585613250732422,\ - 90.66416931152344, 94.88548278808594, 77.90852355957031, 116.56438446044922]), + [0.0, 0.9473684210526315, 0.8421052631578947, 0.7368421052631579, 0.631578947368421, 0.5263157894736842,\ + 0.42105263157894735, 0.3157894736842105, 0.21052631578947367, 0.10526315789473684], + [10.988012313842773, 10.99213981628418, 112.89633178710938, 31.114538192749023, 80.94029235839844,\ + 131.34490966796875, 111.4329605102539, 133.74252319335938, 219.37344360351562, 109.67041015625]), (1.0, optim.lr_scheduler.PolyWarmupLRScheduler,\ - [0.0, 0.2, 0.4, 0.6, 0.8, 1.0, 0.80000002, 0.60000004, 0.40000006000000005, 0.20000007999999997], - [10.988012313842773, 10.99213981628418, 18.11327362060547, 35.164119720458984, 40.68449401855469, 24.585613250732422,\ - 90.66416931152344, 94.88548278808594, 77.90852355957031, 116.56439971923828]) + [0.0, 0.9473684263157895, 0.8421052789473684, 0.7368421315789474, 0.6315789842105263, 0.5263158368421054, + 0.42105268947368424, 0.31578954210526317, 0.21052639473684212, 0.10526324736842106], + [10.988012313842773, 10.99213981628418, 112.89633178710938, 31.114538192749023, 80.9402847290039,\ + 131.3447265625, 111.43253326416016, 133.7415008544922, 219.37147521972656, 109.66986083984375]) ]) def testToyBERTModelLRScheduler(initial_lr, lr_scheduler, expected_learning_rates, expected_losses): # Common setup device = 'cuda' total_steps = 10 seed = 1 + warmup = 0.05 + cycles = 0.5 + power = 1. + lr_end = 1e-7 torch.manual_seed(seed) onnxruntime.set_seed(seed) + # Setup LR Schedulers + if lr_scheduler == optim.lr_scheduler.ConstantWarmupLRScheduler or lr_scheduler == optim.lr_scheduler.LinearWarmupLRScheduler: + lr_scheduler = lr_scheduler(total_steps=total_steps, warmup=warmup) + elif lr_scheduler == optim.lr_scheduler.CosineWarmupLRScheduler: + lr_scheduler = lr_scheduler(total_steps=total_steps, warmup=warmup, cycles=cycles) + elif lr_scheduler == optim.lr_scheduler.PolyWarmupLRScheduler: + lr_scheduler = lr_scheduler(total_steps=total_steps, warmup=warmup, power=power, lr_end=lr_end) + else: + raise RuntimeError("Invalid lr_scheduler") + # Modeling model_desc = bert_model_description() model = load_bert_onnx_model() @@ -248,7 +290,7 @@ def testToyBERTModelLRScheduler(initial_lr, lr_scheduler, expected_learning_rate 'device': { 'id': device, }, - 'lr_scheduler' : lr_scheduler(total_steps=total_steps, warmup=0.5) + 'lr_scheduler' : lr_scheduler }) trainer = orttrainer.ORTTrainer(model, model_desc, optim_config, options=opts) @@ -267,11 +309,11 @@ def testToyBERTModelLRScheduler(initial_lr, lr_scheduler, expected_learning_rate # Dynamic Loss Scaler implemented implicitly @pytest.mark.parametrize("loss_scaler, expected_losses", [ - (None, [10.98803424835205, 10.99240493774414, 11.090575218200684, 11.042827606201172, 10.988829612731934, - 11.105679512023926, 10.981969833374023, 11.08173656463623, 10.997121810913086, 11.10731315612793]), - (amp.DynamicLossScaler(), [10.98803424835205, 10.99240493774414, 11.090575218200684, 11.042827606201172, - 10.988829612731934, 11.105679512023926, 10.981969833374023, 11.081737518310547, 10.99714183807373, 11.107304573059082]), - (CustomLossScaler(), [10.98803424835205, 10.99240493774414, 11.090554237365723, 11.042823791503906, 10.98877239227295, + (None, [10.98803424835205, 10.99240493774414, 11.090575218200684, 11.042827606201172, 10.988829612731934,\ + 11.105679512023926, 10.981968879699707, 11.081787109375, 10.997162818908691, 11.107288360595703]), + (amp.DynamicLossScaler(), [10.98803424835205, 10.99240493774414, 11.090575218200684, 11.042827606201172,\ + 10.988829612731934, 11.105679512023926, 10.981969833374023, 11.081744194030762, 10.997139930725098, 11.107272148132324]), + (CustomLossScaler(), [10.98803424835205, 10.99240493774414, 11.090554237365723, 11.042823791503906, 10.98877239227295,\ 11.105667114257812, 10.981982231140137, 11.081765174865723, 10.997125625610352, 11.107298851013184]) ]) def testToyBERTModelMixedPrecisionLossScaler(loss_scaler, expected_losses): @@ -594,7 +636,7 @@ def testToyBERTSaveAsONNX(): os.remove(onnx_file_name) # Create a new trainer from persisted ONNX model and compare with original ONNX model - trainer_from_onnx = orttrainer.ORTTrainer(reload_onnx_model, model_desc, optim_config)#, options=opts) + trainer_from_onnx = orttrainer.ORTTrainer(reload_onnx_model, model_desc, optim_config, options=opts) assert trainer_from_onnx._onnx_model is not None assert (id(trainer_from_onnx._onnx_model) != id(trainer._onnx_model)) for initializer, loaded_initializer in zip(trainer._onnx_model.graph.initializer, trainer_from_onnx._onnx_model.graph.initializer): @@ -849,3 +891,61 @@ def testToyBERTModelGradientAccumulationLegacyExperimental(gradient_accumulation # Check results _test_helpers.assert_model_outputs(experimental_losses, legacy_losses, rtol=1e-6) + +@pytest.mark.parametrize("params, legacy_optim_map", [ + # Change the hyper parameters for all parameters + ([], legacy_optim_params_a), + # Change the hyperparameters for a subset of hardcoded parameters + ([{'params':['bert.embeddings.LayerNorm.bias', 'bert.embeddings.LayerNorm.weight'], "alpha": 0.9, + "beta": 0.999, "lambda_coef": 0.0, "epsilon": 1e-6, "do_bias_correction":False}], legacy_optim_params_b), + # Change the hyperparameters for a generated set of paramers + (optimizer_parameters(load_bert_onnx_model()), legacy_optim_params_c) +]) +def testToyBERTModelLegacyExperimentalCustomOptimParameters(params, legacy_optim_map): + # Common setup + total_steps = 10 + device = "cuda" + seed = 1 + + # EXPERIMENTAL API + model_desc = bert_model_description() + model = load_bert_onnx_model() + + optim_config = optim.LambConfig(params, alpha= 0.9, beta= 0.999, lambda_coef= 0.01, epsilon= 1e-6, do_bias_correction=False) + opts = orttrainer.ORTTrainerOptions({ + 'debug' : { + 'deterministic_compute': True + }, + 'device': { + 'id': device, + }, + }) + + torch.manual_seed(seed) + onnxruntime.set_seed(seed) + trainer = orttrainer.ORTTrainer(model, model_desc, optim_config, options=opts) + + experimental_losses = [] + for i in range(total_steps): + sample_input = generate_random_input_from_model_desc(model_desc, i) + experimental_losses.append(trainer.train_step(*sample_input).cpu().item()) + + # LEGACY IMPLEMENTATION + device = torch.device(device) + legacy_model_desc, learning_rate_description, learning_rate = legacy_model_params(trainer.optim_config.lr) + torch.manual_seed(seed) + onnxruntime.set_seed(seed) + + legacy_trainer = Legacy_ORTTrainer(model, None, legacy_model_desc, "LambOptimizer", + legacy_optim_map, + learning_rate_description, + device, + _use_deterministic_compute=True) + legacy_losses = [] + for i in range(total_steps): + sample_input = generate_random_input_from_model_desc(model_desc, i) + legacy_sample_input = [*sample_input, learning_rate] + legacy_losses.append(legacy_trainer.train_step(legacy_sample_input).cpu().item()) + + # Check results + _test_helpers.assert_model_outputs(experimental_losses, legacy_losses) diff --git a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py index d2d550468b..faa781b962 100644 --- a/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py +++ b/orttraining/orttraining/test/python/orttraining_test_orttrainer_frontend.py @@ -1,13 +1,12 @@ from functools import partial import inspect import math - +from numpy.testing import assert_allclose import onnx import os import pytest import torch -from numpy.testing import assert_allclose from onnxruntime import set_seed from onnxruntime.capi.ort_trainer import IODescription as Legacy_IODescription,\ @@ -966,8 +965,6 @@ def testORTTrainerLegacyAndExperimentalGradientAccumulation(seed, device, gradie _test_helpers.assert_model_outputs(legacy_loss, experimental_loss, rtol=1e-6) - - @pytest.mark.parametrize("seed,device,optimizer_config,lr_scheduler, get_lr_this_step", [ (0, 'cuda', optim.AdamConfig, optim.lr_scheduler.ConstantWarmupLRScheduler, _test_commons.legacy_constant_lr_scheduler), (0, 'cuda', optim.LambConfig, optim.lr_scheduler.ConstantWarmupLRScheduler, _test_commons.legacy_constant_lr_scheduler), diff --git a/tools/ci_build/build.py b/tools/ci_build/build.py index 2ed0ed87eb..d18e0eab7f 100755 --- a/tools/ci_build/build.py +++ b/tools/ci_build/build.py @@ -1093,8 +1093,6 @@ def run_training_python_frontend_tests(cwd): run_subprocess([ sys.executable, 'orttraining_test_transformers.py', 'BertModelTest.test_for_pretraining_full_precision_list_and_dict_input'], cwd=cwd) - run_subprocess([sys.executable, 'orttraining_test_orttrainer_frontend.py'], cwd=cwd) - run_subprocess([sys.executable, 'orttraining_test_orttrainer_bert_toy_onnx.py'], cwd=cwd) def run_training_python_frontend_e2e_tests(cwd): diff --git a/tools/ci_build/github/linux/docker/scripts/install_deps.sh b/tools/ci_build/github/linux/docker/scripts/install_deps.sh index 0834253d9b..86e461a037 100755 --- a/tools/ci_build/github/linux/docker/scripts/install_deps.sh +++ b/tools/ci_build/github/linux/docker/scripts/install_deps.sh @@ -110,7 +110,7 @@ if [ $DEVICE_TYPE = "Normal" ]; then elif [ $DEVICE_TYPE = "gpu" ]; then ${PYTHON_EXE} -m pip install sympy==1.1.1 if [[ $BUILD_EXTR_PAR = *--enable_training* ]]; then - ${PYTHON_EXE} -m pip install --upgrade --pre torch==1.6.0.dev20200610 torchvision==0.7.0.dev20200610 -f https://download.pytorch.org/whl/nightly/cu101/torch_nightly.html + ${PYTHON_EXE} -m pip install --upgrade --pre torch==1.6.0.dev20200610 torchvision==0.7.0.dev20200610 torchtext==0.6.0.dev20200610 -f https://download.pytorch.org/whl/nightly/cu101/torch_nightly.html ${PYTHON_EXE} -m pip install transformers==v2.10.0 # transformers requires sklearn ${PYTHON_EXE} -m pip install sklearn