diff --git a/mooncake-ep/include/mooncake_worker.cuh b/mooncake-ep/include/mooncake_worker.cuh index 6906ed5a..7315fcde 100644 --- a/mooncake-ep/include/mooncake_worker.cuh +++ b/mooncake-ep/include/mooncake_worker.cuh @@ -34,15 +34,15 @@ __global__ struct Task { void* transferGroupMeta; }; -static constexpr size_t kBufferSize = 1u << 29; +static constexpr size_t kBufferSize = 1u << 24; static constexpr size_t kMaxNumRanks = 64; -void launchReduceKernel(at::Tensor dst, void* src, size_t numRanks, - c10d::ReduceOp op, bool* activeRanks, +void launchReduceKernel(at::Tensor dst, size_t pos, size_t realSize, void* src, + size_t numRanks, c10d::ReduceOp op, bool* activeRanks, cudaStream_t stream); -void launchReduceCpu(at::Tensor dst, void* src, size_t numRanks, - c10d::ReduceOp op); +void launchReduceCpu(at::Tensor dst, size_t pos, size_t realSize, void* src, + size_t numRanks, c10d::ReduceOp op); class MooncakeWorker { public: @@ -51,14 +51,18 @@ class MooncakeWorker { c10::intrusive_ptr putTaskCpu( c10d::OpType opType, size_t tensorSize, int64_t broadcastRoot, TransferGroupMeta* meta, - const std::function& tensorToBuffer, - const std::function& bufferToTensor); + const std::function& + tensorToBuffer, + const std::function& + bufferToTensor); c10::intrusive_ptr putTaskCuda( c10d::OpType opType, size_t tensorSize, int64_t broadcastRoot, TransferGroupMeta* meta, const at::cuda::CUDAStream& stream, - const std::function& tensorToBuffer, - const std::function& bufferToTensor); + const std::function& + tensorToBuffer, + const std::function& + bufferToTensor); void startWorker(); diff --git a/mooncake-ep/src/mooncake_backend.cpp b/mooncake-ep/src/mooncake_backend.cpp index 1b0a80bb..14b7e372 100644 --- a/mooncake-ep/src/mooncake_backend.cpp +++ b/mooncake-ep/src/mooncake_backend.cpp @@ -211,25 +211,27 @@ c10::intrusive_ptr MooncakeBackend::broadcast( if (isCpu_) { return worker_.putTaskCpu( c10d::OpType::BROADCAST, tensorSize, root, &meta_, - [=](void* dst) { + [=](void* dst, size_t pos, size_t realSize) { if (isRoot) { - memcpy(dst, tensor.data_ptr(), tensorSize); + memcpy(dst, (char*)tensor.data_ptr() + pos, realSize); } }, - [=](void* src) { memcpy(tensor.data_ptr(), src, tensorSize); }); + [=](void* src, size_t pos, size_t realSize) { + memcpy((char*)tensor.data_ptr() + pos, src, realSize); + }); } else { at::cuda::CUDAStream stream = at::cuda::getCurrentCUDAStream(tensor.device().index()); return worker_.putTaskCuda( c10d::OpType::BROADCAST, tensorSize, root, &meta_, stream, - [&](void* dst) { + [&](void* dst, size_t pos, size_t realSize) { if (isRoot) { - cudaMemcpyAsync(dst, tensor.data_ptr(), tensorSize, - cudaMemcpyHostToDevice, stream); + cudaMemcpyAsync(dst, (char*)tensor.data_ptr() + pos, + realSize, cudaMemcpyHostToDevice, stream); } }, - [&](void* src) { - cudaMemcpyAsync(tensor.data_ptr(), src, tensorSize, + [&](void* src, size_t pos, size_t realSize) { + cudaMemcpyAsync((char*)tensor.data_ptr() + pos, src, realSize, cudaMemcpyDeviceToHost, stream); }); } @@ -245,23 +247,28 @@ c10::intrusive_ptr MooncakeBackend::allreduce( auto numRanks = size_; return worker_.putTaskCpu( c10d::OpType::ALLREDUCE, tensorSize, 0, &meta_, - [=](void* dst) { memcpy(dst, tensor.data_ptr(), tensorSize); }, - [=](void* src) { - memset(tensor.data_ptr(), 0, tensorSize); - launchReduceCpu(tensor, src, numRanks, opts.reduceOp); + [=](void* dst, size_t pos, size_t realSize) { + memcpy(dst, (char*)tensor.data_ptr() + pos, realSize); + }, + [=](void* src, size_t pos, size_t realSize) { + memset((char*)tensor.data_ptr() + pos, 0, realSize); + launchReduceCpu(tensor, pos, realSize, src, numRanks, + opts.reduceOp); }); } else { auto stream = at::cuda::getCurrentCUDAStream(tensor.device().index()); return worker_.putTaskCuda( c10d::OpType::ALLREDUCE, tensorSize, 0, &meta_, stream, - [&](void* dst) { - cudaMemcpyAsync(dst, tensor.data_ptr(), tensorSize, + [&](void* dst, size_t pos, size_t realSize) { + cudaMemcpyAsync(dst, (char*)tensor.data_ptr() + pos, realSize, cudaMemcpyHostToDevice, stream); }, - [&](void* src) { - cudaMemsetAsync(tensor.data_ptr(), 0, tensorSize, stream); - launchReduceKernel(tensor, src, size_, opts.reduceOp, - meta_.activeRanksDevice, stream); + [&](void* src, size_t pos, size_t realSize) { + cudaMemsetAsync((char*)tensor.data_ptr() + pos, 0, realSize, + stream); + launchReduceKernel(tensor, pos, realSize, src, size_, + opts.reduceOp, meta_.activeRanksDevice, + stream); }); } } @@ -277,11 +284,13 @@ c10::intrusive_ptr MooncakeBackend::allgather( if (isCpu_) { return worker_.putTaskCpu( c10d::OpType::ALLGATHER, tensorSize, 0, &meta_, - [=](void* dst) { memcpy(dst, inputTensor.data_ptr(), tensorSize); }, - [=](void* src) { + [=](void* dst, size_t pos, size_t realSize) { + memcpy(dst, (char*)inputTensor.data_ptr() + pos, realSize); + }, + [=](void* src, size_t pos, size_t realSize) { for (const auto j : c10::irange(outputTensors_.size())) { - memcpy(outputTensors_[j].data_ptr(), src + j * tensorSize, - tensorSize); + memcpy((char*)outputTensors_[j].data_ptr() + pos, + (char*)src + j * realSize, realSize); } }); } else { @@ -289,14 +298,14 @@ c10::intrusive_ptr MooncakeBackend::allgather( at::cuda::getCurrentCUDAStream(inputTensor.device().index()); return worker_.putTaskCuda( c10d::OpType::ALLGATHER, tensorSize, 0, &meta_, stream, - [&](void* dst) { - cudaMemcpyAsync(dst, inputTensor.data_ptr(), tensorSize, - cudaMemcpyHostToDevice, stream); + [&](void* dst, size_t pos, size_t realSize) { + cudaMemcpyAsync(dst, (char*)inputTensor.data_ptr() + pos, + realSize, cudaMemcpyHostToDevice, stream); }, - [&](void* src) { + [&](void* src, size_t pos, size_t realSize) { for (const auto j : c10::irange(outputTensors_.size())) { - cudaMemcpyAsync(outputTensors_[j].data_ptr(), - src + j * tensorSize, tensorSize, + cudaMemcpyAsync((char*)outputTensors_[j].data_ptr() + pos, + (char*)src + j * realSize, realSize, cudaMemcpyDeviceToHost, stream); } }); @@ -311,23 +320,32 @@ c10::intrusive_ptr MooncakeBackend::_allgather_base( auto numRanks = size_; return worker_.putTaskCpu( c10d::OpType::_ALLGATHER_BASE, tensorSize, 0, &meta_, - [=](void* dst) { memcpy(dst, inputBuffer.data_ptr(), tensorSize); }, - [=](void* src) { - memcpy(outputBuffer.data_ptr(), src, tensorSize * numRanks); + [=](void* dst, size_t pos, size_t realSize) { + memcpy(dst, (char*)inputBuffer.data_ptr() + pos, realSize); + }, + [=](void* src, size_t pos, size_t realSize) { + for (const auto j : c10::irange(numRanks)) { + memcpy( + (char*)outputBuffer.data_ptr() + j * tensorSize + pos, + (char*)src + j * realSize, realSize); + } }); } else { auto stream = at::cuda::getCurrentCUDAStream(inputBuffer.device().index()); return worker_.putTaskCuda( c10d::OpType::_ALLGATHER_BASE, tensorSize, 0, &meta_, stream, - [&](void* dst) { - cudaMemcpyAsync(dst, inputBuffer.data_ptr(), tensorSize, - cudaMemcpyHostToDevice, stream); + [&](void* dst, size_t pos, size_t realSize) { + cudaMemcpyAsync(dst, (char*)inputBuffer.data_ptr() + pos, + realSize, cudaMemcpyHostToDevice, stream); }, - [&](void* src) { - cudaMemcpyAsync(outputBuffer.data_ptr(), src, - tensorSize * size_, cudaMemcpyDeviceToHost, - stream); + [&](void* src, size_t pos, size_t realSize) { + for (const auto j : c10::irange(size_)) { + cudaMemcpyAsync( + (char*)outputBuffer.data_ptr() + j * tensorSize + pos, + (char*)src + j * realSize, realSize, + cudaMemcpyDeviceToHost, stream); + } }); } } @@ -340,26 +358,31 @@ c10::intrusive_ptr MooncakeBackend::_reduce_scatter_base( auto numRanks = size_; return worker_.putTaskCpu( c10d::OpType::REDUCE_SCATTER, tensorSize, 0, &meta_, - [=](void* dst) { - memcpy(dst, inputBuffer.data_ptr(), tensorSize * numRanks); + [=](void* dst, size_t pos, size_t realSize) { + memcpy(dst, (char*)inputBuffer.data_ptr() + pos, + realSize * numRanks); }, - [=](void* src) { - memset(outputBuffer.data_ptr(), 0, tensorSize); - launchReduceCpu(outputBuffer, src, numRanks, opts.reduceOp); + [=](void* src, size_t pos, size_t realSize) { + memset((char*)outputBuffer.data_ptr() + pos, 0, realSize); + launchReduceCpu(outputBuffer, pos, realSize, src, numRanks, + opts.reduceOp); }); } else { auto stream = at::cuda::getCurrentCUDAStream(inputBuffer.device().index()); return worker_.putTaskCuda( c10d::OpType::REDUCE_SCATTER, tensorSize, 0, &meta_, stream, - [&](void* dst) { - cudaMemcpyAsync(dst, inputBuffer.data_ptr(), tensorSize * size_, - cudaMemcpyHostToDevice, stream); + [&](void* dst, size_t pos, size_t realSize) { + cudaMemcpyAsync(dst, (char*)inputBuffer.data_ptr() + pos, + realSize * size_, cudaMemcpyHostToDevice, + stream); }, - [&](void* src) { - cudaMemsetAsync(outputBuffer.data_ptr(), 0, tensorSize, stream); - launchReduceKernel(outputBuffer, src, size_, opts.reduceOp, - meta_.activeRanksDevice, stream); + [&](void* src, size_t pos, size_t realSize) { + cudaMemsetAsync((char*)outputBuffer.data_ptr() + pos, 0, + realSize, stream); + launchReduceKernel(outputBuffer, pos, realSize, src, size_, + opts.reduceOp, meta_.activeRanksDevice, + stream); }); } } @@ -372,16 +395,16 @@ c10::intrusive_ptr MooncakeBackend::alltoall( if (isCpu_) { return worker_.putTaskCpu( c10d::OpType::ALLTOALL, tensorSize, 0, &meta_, - [=](void* dst) { + [=](void* dst, size_t pos, size_t realSize) { for (const auto j : c10::irange(inputTensors.size())) { - memcpy(dst + j * tensorSize, inputTensors[j].data_ptr(), - tensorSize); + memcpy(dst + j * realSize, + (char*)inputTensors[j].data_ptr() + pos, realSize); } }, - [=](void* src) { + [=](void* src, size_t pos, size_t realSize) { for (const auto j : c10::irange(outputTensors.size())) { - memcpy(outputTensors[j].data_ptr(), src + j * tensorSize, - tensorSize); + memcpy((char*)outputTensors[j].data_ptr() + pos, + (char*)src + j * realSize, realSize); } }); } else { @@ -389,17 +412,17 @@ c10::intrusive_ptr MooncakeBackend::alltoall( at::cuda::getCurrentCUDAStream(inputTensors[0].device().index()); return worker_.putTaskCuda( c10d::OpType::ALLTOALL, tensorSize, 0, &meta_, stream, - [&](void* dst) { + [&](void* dst, size_t pos, size_t realSize) { for (const auto j : c10::irange(inputTensors.size())) { - cudaMemcpyAsync(dst + j * tensorSize, - inputTensors[j].data_ptr(), tensorSize, - cudaMemcpyHostToDevice, stream); + cudaMemcpyAsync(dst + j * realSize, + (char*)inputTensors[j].data_ptr() + pos, + realSize, cudaMemcpyHostToDevice, stream); } }, - [&](void* src) { + [&](void* src, size_t pos, size_t realSize) { for (const auto j : c10::irange(outputTensors.size())) { - cudaMemcpyAsync(outputTensors[j].data_ptr(), - src + j * tensorSize, tensorSize, + cudaMemcpyAsync((char*)outputTensors[j].data_ptr() + pos, + (char*)src + j * realSize, realSize, cudaMemcpyDeviceToHost, stream); } }); @@ -409,6 +432,7 @@ c10::intrusive_ptr MooncakeBackend::barrier( const c10d::BarrierOptions& opts) { TORCH_CHECK(isCpu_, "Barrier is available only for CPU.") return worker_.putTaskCpu( - c10d::OpType::BARRIER, 0, 0, &meta_, [=](void*) {}, [=](void*) {}); + c10d::OpType::BARRIER, 0, 0, &meta_, [=](void*, size_t, size_t) {}, + [=](void*, size_t, size_t) {}); } } // namespace mooncake diff --git a/mooncake-ep/src/mooncake_worker.cu b/mooncake-ep/src/mooncake_worker.cu index abe8e6f1..91d47d5a 100644 --- a/mooncake-ep/src/mooncake_worker.cu +++ b/mooncake-ep/src/mooncake_worker.cu @@ -78,55 +78,50 @@ __global__ void reduceKernel(scalar_t* dst, const scalar_t* src, } } -void launchReduceKernel(at::Tensor dst, void* src, size_t numRanks, - c10d::ReduceOp op, bool* activeRanks, +void launchReduceKernel(at::Tensor dst, size_t pos, size_t realSize, void* src, + size_t numRanks, c10d::ReduceOp op, bool* activeRanks, cudaStream_t stream) { TORCH_CHECK(op == c10d::ReduceOp::SUM, "Only support SUM for reduction."); + auto ptr = (char*)dst.data_ptr() + pos; + size_t num = realSize / dst.element_size(); + switch (dst.scalar_type()) { case c10::kByte: - reduceKernel<<<64, 256, 0, stream>>>(dst.data_ptr(), - (uint8_t*)src, dst.numel(), - numRanks, activeRanks); + reduceKernel<<<64, 256, 0, stream>>>((uint8_t*)ptr, (uint8_t*)src, + num, numRanks, activeRanks); break; case c10::kChar: - reduceKernel<<<64, 256, 0, stream>>>(dst.data_ptr(), - (int8_t*)src, dst.numel(), - numRanks, activeRanks); + reduceKernel<<<64, 256, 0, stream>>>((int8_t*)ptr, (int8_t*)src, + num, numRanks, activeRanks); break; case c10::kShort: - reduceKernel<<<64, 256, 0, stream>>>(dst.data_ptr(), - (int16_t*)src, dst.numel(), - numRanks, activeRanks); + reduceKernel<<<64, 256, 0, stream>>>((int16_t*)ptr, (int16_t*)src, + num, numRanks, activeRanks); break; case c10::kInt: - reduceKernel<<<64, 256, 0, stream>>>(dst.data_ptr(), (int*)src, - dst.numel(), numRanks, - activeRanks); - break; - case c10::kLong: - reduceKernel<<<64, 256, 0, stream>>>(dst.data_ptr(), - (int64_t*)src, dst.numel(), + reduceKernel<<<64, 256, 0, stream>>>((int*)ptr, (int*)src, num, numRanks, activeRanks); break; + case c10::kLong: + reduceKernel<<<64, 256, 0, stream>>>((int64_t*)ptr, (int64_t*)src, + num, numRanks, activeRanks); + break; case c10::kFloat: - reduceKernel<<<64, 256, 0, stream>>>(dst.data_ptr(), - (float*)src, dst.numel(), + reduceKernel<<<64, 256, 0, stream>>>((float*)ptr, (float*)src, num, numRanks, activeRanks); break; case c10::kDouble: - reduceKernel<<<64, 256, 0, stream>>>(dst.data_ptr(), - (double*)src, dst.numel(), - numRanks, activeRanks); + reduceKernel<<<64, 256, 0, stream>>>((double*)ptr, (double*)src, + num, numRanks, activeRanks); break; case c10::kBool: - reduceKernel<<<64, 256, 0, stream>>>(dst.data_ptr(), - (bool*)src, dst.numel(), + reduceKernel<<<64, 256, 0, stream>>>((bool*)ptr, (bool*)src, num, numRanks, activeRanks); break; case c10::kBFloat16: - reduceKernel<<<64, 256, 0, stream>>>( - dst.data_ptr(), (at::BFloat16*)src, dst.numel(), - numRanks, activeRanks); + reduceKernel<<<64, 256, 0, stream>>>((at::BFloat16*)ptr, + (at::BFloat16*)src, num, + numRanks, activeRanks); break; default: TORCH_CHECK(false, c10::str("Unsupported reduce dtype: ", @@ -164,40 +159,35 @@ void reduceCpu(T* dst, const T* src, size_t numElements, size_t numRanks, }); } -void launchReduceCpu(at::Tensor dst, void* src, size_t numRanks, - c10d::ReduceOp op) { +void launchReduceCpu(at::Tensor dst, size_t pos, size_t realSize, void* src, + size_t numRanks, c10d::ReduceOp op) { + auto ptr = (char*)dst.data_ptr() + pos; + size_t num = realSize / dst.element_size(); + switch (dst.scalar_type()) { case c10::kByte: - reduceCpu(dst.data_ptr(), (uint8_t*)src, dst.numel(), - numRanks, op); + reduceCpu((uint8_t*)ptr, (uint8_t*)src, num, numRanks, op); break; case c10::kChar: - reduceCpu(dst.data_ptr(), (int8_t*)src, dst.numel(), - numRanks, op); + reduceCpu((int8_t*)ptr, (int8_t*)src, num, numRanks, op); break; case c10::kShort: - reduceCpu(dst.data_ptr(), (int16_t*)src, dst.numel(), - numRanks, op); + reduceCpu((int16_t*)ptr, (int16_t*)src, num, numRanks, op); break; case c10::kInt: - reduceCpu(dst.data_ptr(), (int*)src, dst.numel(), numRanks, - op); + reduceCpu((int*)ptr, (int*)src, num, numRanks, op); break; case c10::kLong: - reduceCpu(dst.data_ptr(), (int64_t*)src, dst.numel(), - numRanks, op); + reduceCpu((int64_t*)ptr, (int64_t*)src, num, numRanks, op); break; case c10::kFloat: - reduceCpu(dst.data_ptr(), (float*)src, dst.numel(), numRanks, - op); + reduceCpu((float*)ptr, (float*)src, num, numRanks, op); break; case c10::kDouble: - reduceCpu(dst.data_ptr(), (double*)src, dst.numel(), - numRanks, op); + reduceCpu((double*)ptr, (double*)src, num, numRanks, op); break; case c10::kBool: - reduceCpu(dst.data_ptr(), (bool*)src, dst.numel(), numRanks, - op); + reduceCpu((bool*)ptr, (bool*)src, num, numRanks, op); break; default: TORCH_CHECK(false, c10::str("Unsupported reduce dtype: ", @@ -220,62 +210,106 @@ MooncakeWorker::MooncakeWorker() { c10::intrusive_ptr MooncakeWorker::putTaskCpu( c10d::OpType opType, size_t tensorSize, int64_t broadcastRoot, TransferGroupMeta* meta, - const std::function& tensorToBuffer, - const std::function& bufferToTensor) { - TORCH_CHECK(tensorSize * meta->size < kBufferSize, "Too large!"); + const std::function& + tensorToBuffer, + const std::function& + bufferToTensor) { + size_t chunkSize = ((kBufferSize - 1) / meta->size) & ~(size_t)7; auto future = c10::make_intrusive( c10::ListType::create(c10::TensorType::get())); - // Alternately use even-odd items to maintain tasks - int taskId = cpuTaskCount % 2; - TORCH_CHECK(!tasks_[taskId].active); - int bufferOffset = meta->bufferBaseIndex + meta->taskCount % 2; - tasks_[taskId].opType = opType; - tasks_[taskId].tensorSize = tensorSize; - tasks_[taskId].broadcastRoot = broadcastRoot; - tasks_[taskId].bufferOffset = bufferOffset; - tasks_[taskId].transferGroupMeta = meta; - tensorToBuffer( - (void*)meta->segmentDescs[meta->rank]->buffers[bufferOffset].addr); + struct IterState { + size_t currentPos = 0; + }; + auto state = std::make_shared(); - hasCallback_[taskId] = true; - callbacks_[taskId] = [this, meta, bufferToTensor, bufferOffset, future] { - for (int i = 0; i < meta->size; ++i) { - meta->activeRanksTensor[i] = meta->activeRanks[i] ? 1 : 0; + auto processNextChunk = std::make_shared>(); + + *processNextChunk = [this, processNextChunk, state, opType, tensorSize, + chunkSize, broadcastRoot, meta, tensorToBuffer, + bufferToTensor, future]() { + if (state->currentPos >= tensorSize) { + future->markCompleted(c10::IValue()); + return; } - bufferToTensor((void*)meta->segmentDescs[meta->rank] - ->buffers[bufferOffset + 2] - .addr); - future->markCompleted(c10::IValue()); + + int taskId = cpuTaskCount % 2; + TORCH_CHECK(!tasks_[taskId].active); + + size_t realSize = std::min(chunkSize, tensorSize - state->currentPos); + int bufferOffset = meta->bufferBaseIndex + meta->taskCount % 2; + + tasks_[taskId].opType = opType; + tasks_[taskId].tensorSize = realSize; + tasks_[taskId].broadcastRoot = broadcastRoot; + tasks_[taskId].bufferOffset = bufferOffset; + tasks_[taskId].transferGroupMeta = meta; + + tensorToBuffer( + (void*)meta->segmentDescs[meta->rank]->buffers[bufferOffset].addr, + state->currentPos, realSize); + + hasCallback_[taskId] = true; + + callbacks_[taskId] = [this, processNextChunk, state, meta, + bufferToTensor, bufferOffset, realSize, + future]() { + for (int i = 0; i < meta->size; ++i) { + meta->activeRanksTensor[i] = meta->activeRanks[i] ? 1 : 0; + } + + bufferToTensor((void*)meta->segmentDescs[meta->rank] + ->buffers[bufferOffset + 2] + .addr, + state->currentPos, realSize); + + state->currentPos += realSize; + + (*processNextChunk)(); + }; + + tasks_[taskId].active = true; + ++cpuTaskCount; + ++meta->taskCount; }; - tasks_[taskId].active = true; - ++cpuTaskCount; - ++meta->taskCount; + (*processNextChunk)(); + return c10::make_intrusive(opType, future); } c10::intrusive_ptr MooncakeWorker::putTaskCuda( c10d::OpType opType, size_t tensorSize, int64_t broadcastRoot, TransferGroupMeta* meta, const at::cuda::CUDAStream& stream, - const std::function& tensorToBuffer, - const std::function& bufferToTensor) { - TORCH_CHECK(tensorSize * meta->size < kBufferSize, "Too large!"); - // Alternately use even-odd items to maintain tasks - int taskId = cudaTaskCount % 2 + 2; - int bufferOffset = meta->bufferBaseIndex + meta->taskCount % 2; - tensorToBuffer( - (void*)meta->segmentDescs[meta->rank]->buffers[bufferOffset].addr); + const std::function& + tensorToBuffer, + const std::function& + bufferToTensor) { + // TORCH_CHECK(tensorSize * meta->size < kBufferSize, "Too large!"); + // Alternately use even-odd items to maintain tasks + size_t chunkSize = ((kBufferSize - 1) / meta->size) & ~(size_t)7; + + for (size_t pos = 0; pos < tensorSize; pos += chunkSize) { + size_t realSize = min(tensorSize, pos + chunkSize) - pos; + int taskId = cudaTaskCount % 2 + 2; + int bufferOffset = meta->bufferBaseIndex + meta->taskCount % 2; + tensorToBuffer( + (void*)meta->segmentDescs[meta->rank]->buffers[bufferOffset].addr, + pos, realSize); + + hasCallback_[taskId] = false; + enqueueTaskKernel<<<1, 1, 0, stream>>>( + opType, realSize, broadcastRoot, bufferOffset, meta, tasks_device_, + meta->size, meta->activeRanksDevice, + meta->activeRanksTensor.data_ptr(), taskId); + bufferToTensor((void*)meta->segmentDescs[meta->rank] + ->buffers[bufferOffset + 2] + .addr, + pos, realSize); + ++cudaTaskCount; + ++meta->taskCount; + } - hasCallback_[taskId] = false; - enqueueTaskKernel<<<1, 1, 0, stream>>>( - opType, tensorSize, broadcastRoot, bufferOffset, meta, tasks_device_, - meta->size, meta->activeRanksDevice, - meta->activeRanksTensor.data_ptr(), taskId); - bufferToTensor( - (void*)meta->segmentDescs[meta->rank]->buffers[bufferOffset + 2].addr); - ++cudaTaskCount; - ++meta->taskCount; auto event = std::make_shared(torch::kCUDA); event->record(stream); return c10::make_intrusive(opType, event); diff --git a/mooncake-wheel/tests/test_mooncake_backend_chunk.py b/mooncake-wheel/tests/test_mooncake_backend_chunk.py new file mode 100644 index 00000000..dc8ba272 --- /dev/null +++ b/mooncake-wheel/tests/test_mooncake_backend_chunk.py @@ -0,0 +1,64 @@ +import os +import time +import unittest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +from mooncake import ep + +N = 2 ** 24 + +def worker(rank, world_size, results, collective): + torch.cuda.set_device(rank) + dist.init_process_group( + backend="mooncake-cpu", + rank=rank, + world_size=world_size, + pg_options=ep.MooncakeBackendOptions(torch.zeros((world_size,), dtype=torch.int32, device="cpu")), + ) + + if collective == "all_reduce": + tensor = torch.tensor([rank + 1] * N, dtype=torch.int32, device="cpu") + dist.all_reduce(tensor, op=dist.ReduceOp.SUM) + results[rank] = tensor[0].item() + assert torch.all(tensor == tensor[0].item()) + + else: + raise ValueError(f"Unsupported collective: {collective}") + + while len(results) < world_size: + time.sleep(1) + + dist.destroy_process_group() + + +class TestMooncakeBackend(unittest.TestCase): + def setUp(self): + self.world_size = torch.cuda.device_count() + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = "29500" + + def tearDown(self): + pass + + def _spawn_and_check(self, collective): + mp_manager = mp.Manager() + results = mp_manager.dict() + mp.spawn( + worker, + args=(self.world_size, results, collective), + nprocs=self.world_size, + join=True, + ) + + expected = sum(range(1, self.world_size + 1)) + for r in range(self.world_size): + self.assertEqual(results[r], expected) + + def test_allreduce_sum(self): + self._spawn_and_check("all_reduce") + + + +if __name__ == "__main__": + unittest.main()