mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-23 19:32:23 +00:00
MPI Kernels part 2
This commit is contained in:
parent
66babe4ed3
commit
3b2bd5f50d
6 changed files with 16 additions and 14 deletions
|
|
@ -460,6 +460,7 @@ struct ProviderHostImpl : ProviderHost {
|
|||
void KernelDefBuilder__Alias(KernelDefBuilder* p, const std::vector<std::pair<int, int>>& aliases) override { p->Alias(aliases); }
|
||||
void KernelDefBuilder__VariadicAlias(KernelDefBuilder* p, int input_offset, int output_offset) override { p->VariadicAlias(input_offset, output_offset); }
|
||||
void KernelDefBuilder__ExternalOutputs(KernelDefBuilder* p) override { p->ExternalOutputs(); }
|
||||
void KernelDefBuilder__AllocateInputsContiguously(KernelDefBuilder* p) override { p->AllocateInputsContiguously(); }
|
||||
|
||||
std::unique_ptr<KernelDef> KernelDefBuilder__Build(KernelDefBuilder* p) override { return p->Build(); }
|
||||
|
||||
|
|
|
|||
|
|
@ -387,6 +387,7 @@ struct ProviderHost {
|
|||
virtual void KernelDefBuilder__Alias(KernelDefBuilder* p, const std::vector<std::pair<int, int>>& aliases) = 0;
|
||||
virtual void KernelDefBuilder__VariadicAlias(KernelDefBuilder* p, int input_offset, int output_offset) = 0;
|
||||
virtual void KernelDefBuilder__ExternalOutputs(KernelDefBuilder* p) = 0;
|
||||
virtual void KernelDefBuilder__AllocateInputsContiguously(KernelDefBuilder* p) = 0;
|
||||
|
||||
virtual std::unique_ptr<KernelDef> KernelDefBuilder__Build(KernelDefBuilder* p) = 0;
|
||||
|
||||
|
|
@ -1174,6 +1175,11 @@ struct KernelDefBuilder final {
|
|||
return *this;
|
||||
}
|
||||
|
||||
KernelDefBuilder& AllocateInputsContiguously() {
|
||||
g_host->KernelDefBuilder__AllocateInputsContiguously(this);
|
||||
return *this;
|
||||
}
|
||||
|
||||
std::unique_ptr<KernelDef> Build() {
|
||||
return g_host->KernelDefBuilder__Build(this);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,9 +8,9 @@
|
|||
#ifndef SHARED_PROVIDER
|
||||
#include "core/common/logging/logging.h"
|
||||
#include "core/framework/framework_common.h"
|
||||
#include "core/framework/ml_value.h"
|
||||
#include "core/graph/node_arg.h"
|
||||
#endif
|
||||
#include "core/framework/ml_value.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace training {
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "nccl_kernels.h"
|
||||
#include "core/providers/cuda/tensor/identity_op.h"
|
||||
|
||||
|
|
@ -12,7 +11,7 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
kMSDomain,
|
||||
1,
|
||||
kCudaExecutionProvider,
|
||||
KernelDefBuilder()
|
||||
(*KernelDefBuilder::Create())
|
||||
.Alias(0, 0)
|
||||
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes()),
|
||||
IdentityOp<false>);
|
||||
|
|
@ -22,7 +21,7 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
kMSDomain,
|
||||
1,
|
||||
kCudaExecutionProvider,
|
||||
KernelDefBuilder()
|
||||
(*KernelDefBuilder::Create())
|
||||
.Alias(0, 0)
|
||||
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes()),
|
||||
NcclAllReduce);
|
||||
|
|
|
|||
|
|
@ -2,11 +2,9 @@
|
|||
// Licensed under the MIT License.
|
||||
|
||||
#include "orttraining/training_ops/cuda/collective/nccl_common.h"
|
||||
#include "core/common/logging/logging.h"
|
||||
#include "orttraining/core/framework/communication/mpi/mpi_context.h"
|
||||
#include <mpi.h>
|
||||
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace cuda {
|
||||
|
||||
|
|
@ -80,7 +78,7 @@ NcclContext::NcclContext() {
|
|||
|
||||
// Initialize Data Parallel Group NCCL Communicator
|
||||
ret = CreateNcclCommunicator(&mpi_world_group, training::WorkerGroupType::DataParallel,
|
||||
&data_group_comm_);
|
||||
&data_group_comm_);
|
||||
ORT_ENFORCE(ret.IsOK());
|
||||
|
||||
// Initialize Horizontal Model Parallel Group NCCL Communicator
|
||||
|
|
@ -111,11 +109,9 @@ ncclComm_t NcclContext::Comm(training::WorkerGroupType group_type) {
|
|||
return data_group_comm_;
|
||||
} else if (training::WorkerGroupType::HorizontalParallel == group_type) {
|
||||
return horizontal_group_comm_;
|
||||
}
|
||||
else if (training::WorkerGroupType::NodeLocalDataParallel == group_type) {
|
||||
} else if (training::WorkerGroupType::NodeLocalDataParallel == group_type) {
|
||||
return node_local_comm_;
|
||||
}
|
||||
else if (training::WorkerGroupType::CrossNodeDataParallel == group_type) {
|
||||
} else if (training::WorkerGroupType::CrossNodeDataParallel == group_type) {
|
||||
return cross_node_comm_;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -217,7 +217,7 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
kMSDomain,
|
||||
1,
|
||||
kCudaExecutionProvider,
|
||||
KernelDefBuilder()
|
||||
(*KernelDefBuilder::Create())
|
||||
.VariadicAlias(0, 0) // outputs and inputs are mapped one to one
|
||||
.AllocateInputsContiguously()
|
||||
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes()),
|
||||
|
|
@ -228,7 +228,7 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
kMSDomain,
|
||||
1,
|
||||
kCudaExecutionProvider,
|
||||
KernelDefBuilder()
|
||||
(*KernelDefBuilder::Create())
|
||||
.VariadicAlias(0, 0) // outputs and inputs are mapped one to one
|
||||
.AllocateInputsContiguously()
|
||||
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes()),
|
||||
|
|
@ -239,7 +239,7 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
kMSDomain,
|
||||
1,
|
||||
kCudaExecutionProvider,
|
||||
KernelDefBuilder()
|
||||
(*KernelDefBuilder::Create())
|
||||
.VariadicAlias(0, 0) // outputs and inputs are mapped one to one
|
||||
.AllocateInputsContiguously()
|
||||
.TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes()),
|
||||
|
|
|
|||
Loading…
Reference in a new issue