MPI Kernels part 2

This commit is contained in:
Ryan Hill 2021-05-11 21:33:28 -07:00
parent 66babe4ed3
commit 3b2bd5f50d
6 changed files with 16 additions and 14 deletions

View file

@ -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(); }

View file

@ -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);
}

View file

@ -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 {

View file

@ -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);

View file

@ -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_;
}

View file

@ -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()),