mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Test moving DistributedRunContext instance into shared provider layer
(with purpose error to verify it's being built properly)
This commit is contained in:
parent
741e09a882
commit
0a59bc3902
3 changed files with 15 additions and 4 deletions
|
|
@ -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<int64_t>& new_shape, std::vector<size_t>& 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<int64_t>& 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_;
|
||||
|
|
|
|||
|
|
@ -46,6 +46,10 @@ class PassThrough;
|
|||
class YieldOp;
|
||||
} // namespace contrib
|
||||
|
||||
namespace training {
|
||||
class DistributedRunContext;
|
||||
}
|
||||
|
||||
template <typename T, typename TResult>
|
||||
struct IteratorHolder {
|
||||
IteratorHolder(std::unique_ptr<T>&& 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<int64_t>& new_shape, std::vector<size_t>& 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<int64_t>& split_sizes) = 0;
|
||||
virtual Status contrib__YieldOp__Compute(const contrib::YieldOp* p, OpKernelContext* context) = 0;
|
||||
|
||||
virtual training::DistributedRunContext& GetDistributedRunContextInstance() = 0;
|
||||
#endif
|
||||
#endif
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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<int32_t> GetRanks(WorkerGroupType group_type){
|
||||
static std::vector<int32_t> GetRanks(WorkerGroupType group_type) {
|
||||
return DistributedRunContext::GetInstance().GetWorkerGroup(group_type).ranks;
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue