This is an automated email from the ASF dual-hosted git repository.

hongyij pushed a commit to branch unity
in repository https://gitbox.apache.org/repos/asf/tvm.git

commit 1ab89c07e423b20e8807d82f9435adf0c6bdaa6b
Author: Hongyi Jin <[email protected]>
AuthorDate: Thu Sep 14 13:39:18 2023 -0400

    change stream
---
 src/runtime/disco/nccl/nccl.cc | 18 ++++++++++++++++--
 1 file changed, 16 insertions(+), 2 deletions(-)

diff --git a/src/runtime/disco/nccl/nccl.cc b/src/runtime/disco/nccl/nccl.cc
index e404e3c2bb..7ac8341eca 100644
--- a/src/runtime/disco/nccl/nccl.cc
+++ b/src/runtime/disco/nccl/nccl.cc
@@ -76,7 +76,6 @@ void InitCCLPerWorker(ShapeTuple device_ids, std::string 
unique_id_bytes) {
   CUDA_CALL(cudaSetDevice(device_id));
   CUDA_CALL(cudaStreamCreate(&ctx->stream));
   Device device{DLDeviceType::kDLCUDA, device_id};
-  DeviceAPI::Get(device)->SetStream(device, ctx->stream);
   worker->default_device = device;
   worker->ccl = "nccl";
   ctx->worker = worker;
@@ -91,9 +90,12 @@ void AllReduce(NDArray send, ReduceKind reduce_kind, NDArray 
recv) {
   NCCLThreadLocalContext* ctx = NCCLThreadLocalContext::Get();
   ShapeTuple shape = send.Shape();
   int64_t numel = shape->Product();
+  Device device = ctx->worker->default_device;
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, nullptr, ctx->stream);
   NCCL_CALL(ncclAllReduce(send->data, recv->data, numel,
                           /*datatype=*/AsNCCLDataType(DataType(send->dtype)),
                           /*op=*/AsNCCLRedOp(reduce_kind), ctx->comm, 
ctx->stream));
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, ctx->stream, nullptr);
 }
 
 void BroadcastFromWorker0(NDArray send, NDArray recv) {
@@ -101,9 +103,12 @@ void BroadcastFromWorker0(NDArray send, NDArray recv) {
   ICHECK(send.Shape()->Product() == recv.Shape()->Product());
   ShapeTuple shape = send.Shape();
   int64_t numel = shape->Product();
+  Device device = ctx->worker->default_device;
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, nullptr, ctx->stream);
   NCCL_CALL(ncclBroadcast(send->data, recv->data, numel,
                           /*datatype=*/AsNCCLDataType(DataType(send->dtype)),
                           /*root=*/0, ctx->comm, ctx->stream));
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, ctx->stream, nullptr);
 }
 
 void ScatterFromWorker0(Optional<NDArray> send, NDArray recv) {
@@ -111,6 +116,8 @@ void ScatterFromWorker0(Optional<NDArray> send, NDArray 
recv) {
   NCCLThreadLocalContext* ctx = NCCLThreadLocalContext::Get();
   int worker_id = ctx->worker->worker_id;
   int num_workers = ctx->worker->num_workers;
+  Device device = ctx->worker->default_device;
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, nullptr, ctx->stream);
   if (worker_id == 0) {
     CHECK(send.defined()) << "ValueError: buffer `send` must be provided when 
worker_id == 0.";
     NDArray buffer = send.value();
@@ -144,6 +151,7 @@ void ScatterFromWorker0(Optional<NDArray> send, NDArray 
recv) {
   DataType dtype(recv->dtype);
   NCCL_CALL(ncclRecv(recv->data, numel, AsNCCLDataType(dtype), 0, ctx->comm, 
ctx->stream));
   NCCL_CALL(ncclGroupEnd());
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, ctx->stream, nullptr);
 }
 
 void GatherToWorker0(NDArray send, Optional<NDArray> recv) {
@@ -151,6 +159,8 @@ void GatherToWorker0(NDArray send, Optional<NDArray> recv) {
   NCCLThreadLocalContext* ctx = NCCLThreadLocalContext::Get();
   int worker_id = ctx->worker->worker_id;
   int num_workers = ctx->worker->num_workers;
+  Device device = ctx->worker->default_device;
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, nullptr, ctx->stream);
   if (worker_id == 0) {
     CHECK(recv.defined()) << "ValueError: buffer `recv` must be provided when 
worker_id == 0.";
     NDArray buffer = recv.value();
@@ -184,22 +194,26 @@ void GatherToWorker0(NDArray send, Optional<NDArray> 
recv) {
   DataType dtype(send->dtype);
   NCCL_CALL(ncclSend(send->data, numel, AsNCCLDataType(dtype), 0, ctx->comm, 
ctx->stream));
   NCCL_CALL(ncclGroupEnd());
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, ctx->stream, nullptr);
 }
 
 void RecvFromWorker0(NDArray buffer) {
   NCCLThreadLocalContext* ctx = NCCLThreadLocalContext::Get();
   CHECK_NE(ctx->worker->worker_id, 0)
       << "ValueError: Worker 0 is not allowed to call RecvFromWorker0.";
+  Device device = ctx->worker->default_device;
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, nullptr, ctx->stream);
   NCCL_CALL(ncclGroupStart());
   NCCL_CALL(ncclRecv(buffer->data, buffer.Shape()->Product(), 
AsNCCLDataType(buffer.DataType()), 0,
                      ctx->comm, ctx->stream));
   NCCL_CALL(ncclGroupEnd());
+  DeviceAPI::Get(device)->SyncStreamFromTo(device, ctx->stream, nullptr);
 }
 
 void SyncWorker() {
   NCCLThreadLocalContext* ctx = NCCLThreadLocalContext::Get();
   ICHECK(ctx->worker != nullptr);
-  CUDA_CALL(cudaStreamSynchronize(ctx->stream));
+  CUDA_CALL(cudaStreamSynchronize(nullptr));
 }
 
 TVM_REGISTER_GLOBAL("runtime.disco.nccl.init_ccl").set_body_typed(InitCCL);

Reply via email to