From 3b2bd5f50d24126dbe4899e76a4fe1ab67fed19b Mon Sep 17 00:00:00 2001 From: Ryan Hill Date: Tue, 11 May 2021 21:33:28 -0700 Subject: [PATCH] MPI Kernels part 2 --- onnxruntime/core/framework/provider_bridge_ort.cc | 1 + .../providers/shared_library/provider_interfaces.h | 6 ++++++ orttraining/orttraining/core/graph/optimizer_config.h | 2 +- .../training_ops/cuda/collective/megatron.cc | 5 ++--- .../training_ops/cuda/collective/nccl_common.cc | 10 +++------- .../training_ops/cuda/collective/nccl_kernels.cc | 6 +++--- 6 files changed, 16 insertions(+), 14 deletions(-) diff --git a/onnxruntime/core/framework/provider_bridge_ort.cc b/onnxruntime/core/framework/provider_bridge_ort.cc index e828a385d5..16cbb68883 100644 --- a/onnxruntime/core/framework/provider_bridge_ort.cc +++ b/onnxruntime/core/framework/provider_bridge_ort.cc @@ -460,6 +460,7 @@ struct ProviderHostImpl : ProviderHost { void KernelDefBuilder__Alias(KernelDefBuilder* p, const std::vector>& 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 KernelDefBuilder__Build(KernelDefBuilder* p) override { return p->Build(); } diff --git a/onnxruntime/core/providers/shared_library/provider_interfaces.h b/onnxruntime/core/providers/shared_library/provider_interfaces.h index 0b94225d1a..a46ec21d1a 100644 --- a/onnxruntime/core/providers/shared_library/provider_interfaces.h +++ b/onnxruntime/core/providers/shared_library/provider_interfaces.h @@ -387,6 +387,7 @@ struct ProviderHost { virtual void KernelDefBuilder__Alias(KernelDefBuilder* p, const std::vector>& 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 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 Build() { return g_host->KernelDefBuilder__Build(this); } diff --git a/orttraining/orttraining/core/graph/optimizer_config.h b/orttraining/orttraining/core/graph/optimizer_config.h index 089b84d4ce..934c431e13 100644 --- a/orttraining/orttraining/core/graph/optimizer_config.h +++ b/orttraining/orttraining/core/graph/optimizer_config.h @@ -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 { diff --git a/orttraining/orttraining/training_ops/cuda/collective/megatron.cc b/orttraining/orttraining/training_ops/cuda/collective/megatron.cc index 55ba247966..27bf09c3eb 100644 --- a/orttraining/orttraining/training_ops/cuda/collective/megatron.cc +++ b/orttraining/orttraining/training_ops/cuda/collective/megatron.cc @@ -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); @@ -22,7 +21,7 @@ ONNX_OPERATOR_KERNEL_EX( kMSDomain, 1, kCudaExecutionProvider, - KernelDefBuilder() + (*KernelDefBuilder::Create()) .Alias(0, 0) .TypeConstraint("T", DataTypeImpl::AllIEEEFloatTensorTypes()), NcclAllReduce); diff --git a/orttraining/orttraining/training_ops/cuda/collective/nccl_common.cc b/orttraining/orttraining/training_ops/cuda/collective/nccl_common.cc index aec65686c9..930a1a6937 100644 --- a/orttraining/orttraining/training_ops/cuda/collective/nccl_common.cc +++ b/orttraining/orttraining/training_ops/cuda/collective/nccl_common.cc @@ -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 - 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_; } diff --git a/orttraining/orttraining/training_ops/cuda/collective/nccl_kernels.cc b/orttraining/orttraining/training_ops/cuda/collective/nccl_kernels.cc index 7bd7dabdbc..0db6629f45 100644 --- a/orttraining/orttraining/training_ops/cuda/collective/nccl_kernels.cc +++ b/orttraining/orttraining/training_ops/cuda/collective/nccl_kernels.cc @@ -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()),