diff --git a/onnxruntime/core/framework/provider_bridge_ort.cc b/onnxruntime/core/framework/provider_bridge_ort.cc index 00208d1e41..fff9daa65d 100644 --- a/onnxruntime/core/framework/provider_bridge_ort.cc +++ b/onnxruntime/core/framework/provider_bridge_ort.cc @@ -62,6 +62,7 @@ Status LongformerAttentionBase__CheckInputs(const LongformerAttentionBase* p, co #include "orttraining/training_ops/cpu/controlflow/yield.h" #include "orttraining/training_ops/cpu/loss/softmax_cross_entropy_loss.h" #include "orttraining/training_ops/cpu/tensor/split.h" +#include "orttraining/core/framework/distributed_run_context.h" #endif #if defined(USE_CUDA) && defined(ORT_USE_NCCL) && defined(USE_NCCL_P2P) #include "orttraining/training_ops/cuda/communication/nccl_service.h" @@ -879,6 +880,8 @@ struct ProviderHostImpl : ProviderHost { void contrib__GetPermutationAndShape(bool ncd_to_ndc, const TensorShape& tensor_shape, std::vector& new_shape, std::vector& permutations) override { contrib::GetPermutationAndShape(ncd_to_ndc, tensor_shape, new_shape, permutations); } Status contrib__PrepareForTrainingCompute(const TensorShape& input_shape, int num_outputs, int64_t& axis, int& before_dims, int& after_dims_including_split_axis, int& after_dims_excluding_split, std::vector& split_sizes) override { return contrib::PrepareForTrainingCompute(input_shape, num_outputs, axis, before_dims, after_dims_including_split_axis, after_dims_excluding_split, split_sizes); } Status contrib__YieldOp__Compute(const contrib::YieldOp* p, OpKernelContext* context) override { return p->YieldOp::Compute(context); } + + training::DistributedRunContext& GetDistributedRunContextInstance() override { return training::DistributedRunContext::GetInstance(); } #endif #endif } provider_host_; diff --git a/onnxruntime/core/providers/shared_library/provider_interfaces.h b/onnxruntime/core/providers/shared_library/provider_interfaces.h index 38d91e36a2..3644b4d918 100644 --- a/onnxruntime/core/providers/shared_library/provider_interfaces.h +++ b/onnxruntime/core/providers/shared_library/provider_interfaces.h @@ -46,6 +46,10 @@ class PassThrough; class YieldOp; } // namespace contrib +namespace training { +class DistributedRunContext; +} + template struct IteratorHolder { IteratorHolder(std::unique_ptr&& p) : p_{std::move(p)} {} @@ -784,6 +788,8 @@ struct ProviderHost { virtual void contrib__GetPermutationAndShape(bool ncd_to_ndc, const TensorShape& tensor_shape, std::vector& new_shape, std::vector& permutations) = 0; virtual Status contrib__PrepareForTrainingCompute(const TensorShape& input_shape, int num_outputs, int64_t& axis, int& before_dims, int& after_dims_including_split_axis, int& after_dims_excluding_split, std::vector& split_sizes) = 0; virtual Status contrib__YieldOp__Compute(const contrib::YieldOp* p, OpKernelContext* context) = 0; + + virtual training::DistributedRunContext& GetDistributedRunContextInstance() = 0; #endif #endif }; diff --git a/orttraining/orttraining/core/framework/distributed_run_context.h b/orttraining/orttraining/core/framework/distributed_run_context.h index 366c15523f..436089e5cc 100644 --- a/orttraining/orttraining/core/framework/distributed_run_context.h +++ b/orttraining/orttraining/core/framework/distributed_run_context.h @@ -33,8 +33,7 @@ struct WorkerGroup { std::string ToString() const { std::stringstream msg; - msg << "group_type: " << group_type << ", group_id: " << group_id << - ", rank in group:" << rank_in_group << ", world-rank:" << ranks.at(rank_in_group); + msg << "group_type: " << group_type << ", group_id: " << group_id << ", rank in group:" << rank_in_group << ", world-rank:" << ranks.at(rank_in_group); msg << ", ranks: ["; for (size_t i = 0; i < ranks.size(); ++i) { msg << ranks.at(i); @@ -81,10 +80,13 @@ class DistributedRunContext { config.pipeline_stage_size); } +#ifndef SHARED_PROVIDER static DistributedRunContext& GetInstance() { return DistributedRunContext::GetOrCreateInstance(); } - +#else + DistributedRunContext& GetInstance() { return Provider_Gethost()->GetDistributedRunContextInstance; } +#endif /* SHORTCUT FUNCTIONS START */ static DistributedRunConfig& RunConfig() { @@ -121,7 +123,7 @@ class DistributedRunContext { return DistributedRunContext::GetInstance().GetWorkerGroup(group_type).group_id; } - static std::vector GetRanks(WorkerGroupType group_type){ + static std::vector GetRanks(WorkerGroupType group_type) { return DistributedRunContext::GetInstance().GetWorkerGroup(group_type).ranks; }