[EP] Implement elastic scaling up (#1173)

This commit is contained in:
Xun Sun 2025-12-22 14:10:40 +08:00 committed by GitHub
parent 4c2b6a95f6
commit 832ae19492
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
9 changed files with 533 additions and 139 deletions

View File

@ -255,6 +255,7 @@ jobs:
run: |
source test_env/bin/activate
python -m unittest mooncake-wheel.tests.test_mooncake_backend_cpu
python -m unittest mooncake-wheel.tests.test_mooncake_backend_elastic
shell: bash
test-sglang-integration:

View File

@ -4,7 +4,7 @@
Mooncake EP is an adaption of [DeepEP](https://github.com/deepseek-ai/DeepEP) that supports **fault tolerance** and fast data transfer with **IBGDA**, designed as a critical component for large-scale, latency-sensitive MoE (Mixture of Experts) inference. Mooncake EP aims to retain full compatibility with the DeepEP API, with the addition of an `active_ranks` tensor passed to both the `dispatch` and `combine` functions to capture information about rank activeness. By integrating with the EPLB module, Mooncake EP ensures fault tolerance during MoE inference, enabling robust performance even in large-scale, fault-prone environments.
Mooncake Backend is a PyTorch distributed backend (a replacement for NCCL and Gloo) that provides **fault-tolerant collective communication primitives** and can be seamlessly integrated into machine learning systems. Built with the [Transfer Engine](transfer-engine/index.md), Mooncake Backend ensures that collective communications can continue even in the event of rank failures. Furthermore, it reports these failures to the upper layers of the system, allowing for graceful error handling without disrupting ongoing operations.
Mooncake Backend is a PyTorch distributed backend (a replacement for NCCL and Gloo) that provides **fault-tolerant collective communication primitives** and can be seamlessly integrated into machine learning systems. Built with the [Transfer Engine](../design/transfer-engine/index.md), Mooncake Backend ensures that collective communications can continue even in the event of rank failures. Furthermore, it reports these failures to the upper layers of the system, allowing for graceful error handling without disrupting ongoing operations.
## Usage
@ -68,3 +68,40 @@ assert active_ranks.all() # Verify that no ranks are broken
For a full example, see `mooncake-wheel/tests/test_mooncake_backend.py`.
---
Recover usage (e.g., wants to recover rank #2):
```python
# For the healthy processes, execute:
import torch
import torch.distributed as dist
from mooncake import ep
...
broken_rank = 2
backend = dist.group.WORLD._get_backend(torch.device("cpu"))
while True:
(peer_state,) = ep.get_peer_state(backend, [broken_rank])
if peer_state:
ep.recover_ranks(backend, [broken_rank])
break
else:
# Handle ongoing logic, like inference
pass
# For the new process, execute:
dist.init_process_group(
backend="mooncake-cpu",
rank=broken_rank,
world_size=num_processes,
pg_options=ep.MooncakeBackendOptions(
torch.ones((num_processes,), dtype=torch.int32),
is_extension=True, # Must set this option to True
),
)
```
For a full example, see `mooncake-wheel/tests/test_mooncake_backend_elastic.py`.

View File

@ -13,10 +13,15 @@ class MooncakeBackend final : public ::c10d::Backend {
struct MooncakeBackendOptions final : ::c10d::Backend::Options {
explicit MooncakeBackendOptions(at::Tensor activeRanks)
: Options{"mooncake"}, activeRanks_{activeRanks} {}
MooncakeBackendOptions(at::Tensor activeRanks, bool isExtension)
: Options{"mooncake"},
activeRanks_{activeRanks},
isExtension_{isExtension} {}
~MooncakeBackendOptions() override = default;
at::Tensor activeRanks_;
bool isExtension_ = false;
};
MooncakeBackend(c10::intrusive_ptr<::c10d::Store> store, int rank, int size,
@ -71,9 +76,17 @@ class MooncakeBackend final : public ::c10d::Backend {
at::Tensor getActiveRanksTensor() { return meta_.activeRanksTensor; }
int getNumSyncedRanks();
void extendGroupSizeTo(int size);
std::vector<bool> getPeerState(const std::vector<int>& ranks);
void recoverRanks(const std::vector<int>& ranks);
private:
static TransferEngine engine_;
static Transport* transport_;
static bool engineInitialized_;
static int backendIndex_;
bool isCpu_{false};
static std::string hostIp_;
@ -81,8 +94,15 @@ class MooncakeBackend final : public ::c10d::Backend {
void* recv_buffer_[2];
int32_t* cpu_sync_send_region_[2];
int32_t* cpu_sync_recv_region_[2];
int32_t* warmup_send_region_;
int32_t* warmup_recv_region_;
static MooncakeWorker worker_;
TransferGroupMeta meta_;
bool isShutdown_{false};
int nextRankForConnection_ = 0;
void connectionPoller(c10::intrusive_ptr<::c10d::Store> store,
int backendIndex);
};
} // namespace mooncake

View File

@ -11,6 +11,9 @@
namespace mooncake {
static constexpr size_t kBufferSize = 1u << 24;
static constexpr size_t kMaxNumRanks = 64;
struct TransferGroupMeta {
int rank;
int size;
@ -18,10 +21,13 @@ struct TransferGroupMeta {
bool* activeRanks;
bool* activeRanksDevice;
at::Tensor activeRanksTensor;
bool peerConnected[kMaxNumRanks]{};
TransferEngine* engine;
c10::intrusive_ptr<::c10d::Store> store;
int bufferBaseIndex;
std::vector<TransferMetadata::SegmentID> segmentIDs;
std::vector<std::shared_ptr<TransferMetadata::SegmentDesc>> segmentDescs;
int backendIndex;
TransferMetadata::SegmentID segmentIDs[kMaxNumRanks];
std::shared_ptr<TransferMetadata::SegmentDesc> segmentDescs[kMaxNumRanks];
};
__global__ struct Task {
@ -34,15 +40,12 @@ __global__ struct Task {
void* transferGroupMeta;
};
static constexpr size_t kBufferSize = 1u << 24;
static constexpr size_t kMaxNumRanks = 64;
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, size_t pos, size_t realSize, void* src,
size_t numRanks, c10d::ReduceOp op);
size_t numRanks, c10d::ReduceOp op, bool* activeRanks);
class MooncakeWorker {
public:

View File

@ -18,7 +18,7 @@ constexpr int kBarrierDummyTensorSize = 1;
std::string MooncakeBackend::hostIp_ = "127.0.0.1";
TransferEngine MooncakeBackend::engine_ = TransferEngine(true);
Transport* MooncakeBackend::transport_ = nullptr;
bool MooncakeBackend::engineInitialized_ = false;
int MooncakeBackend::backendIndex_ = 0;
MooncakeWorker MooncakeBackend::worker_;
@ -27,19 +27,11 @@ MooncakeBackend::MooncakeBackend(
c10::intrusive_ptr<MooncakeBackendOptions> options, bool isCpu)
: Backend(rank, size), isCpu_(isCpu) {
// Initialize transfer engine
if (!transport_) {
if (!engineInitialized_) {
engine_.init(P2PHANDSHAKE, hostIp_);
transport_ = engine_.installTransport("rdma", nullptr);
if (!transport_) {
// Fallback to TCP
transport_ = engine_.installTransport("tcp", nullptr);
TORCH_CHECK(transport_ != nullptr,
c10::str("Failed to install transport"));
LOG(WARNING) << "[Torch Backend] RDMA transport unavailable. "
"Fallback to TCP.";
}
engineInitialized_ = true;
}
auto localRpcMeta = transport_->meta()->localRpcMeta();
auto localRpcMeta = engine_.getMetadata()->localRpcMeta();
std::string localServerName = localRpcMeta.ip_or_host_name + ":" +
std::to_string(localRpcMeta.rpc_port);
@ -102,23 +94,27 @@ MooncakeBackend::MooncakeBackend(
TORCH_CHECK(!rc, REGISTER_BUFFER_ERROR_MSG);
}
// Reset the synchronization signal
store->deleteKey("backend_init_" + std::to_string(backendIndex_) + "_" +
std::to_string(rank_));
store->deleteKey("backend_warmup_" + std::to_string(backendIndex_) + "_" +
std::to_string(rank_));
warmup_send_region_ = new int32_t[kMaxNumRanks];
warmup_send_region_[0] = 1;
int rc = engine_.registerLocalMemory(
warmup_send_region_, kMaxNumRanks * sizeof(int32_t), kWildcardLocation);
TORCH_CHECK(!rc, REGISTER_BUFFER_ERROR_MSG);
warmup_recv_region_ = new int32_t[kMaxNumRanks]{};
rc = engine_.registerLocalMemory(
warmup_recv_region_, kMaxNumRanks * sizeof(int32_t), kWildcardLocation);
TORCH_CHECK(!rc, REGISTER_BUFFER_ERROR_MSG);
// Sync metadata
store->set("server_name_" + std::to_string(backendIndex_) + "_" +
std::to_string(rank_),
localServerName);
std::vector<std::string> server_names;
for (int i = 0; i < size; i++) {
server_names.push_back(store->get_to_str("server_name_" +
std::to_string(backendIndex_) +
"_" + std::to_string(i)));
}
int backendIndex = backendIndex_;
std::thread([this, store, backendIndex] {
connectionPoller(store, backendIndex);
}).detach();
meta_.rank = rank;
meta_.size = size;
@ -151,70 +147,25 @@ MooncakeBackend::MooncakeBackend(
.device(isCpu ? torch::kCPU : torch::kCUDA));
}
meta_.engine = &engine_;
meta_.bufferBaseIndex = backendIndex_ * 8;
meta_.segmentIDs.clear();
meta_.segmentDescs.clear();
for (int i = 0; i < size_; ++i) {
auto segment_id = engine_.openSegment(server_names[i]);
meta_.segmentIDs.emplace_back(segment_id);
auto segment_desc =
engine_.getMetadata()->getSegmentDescByID(segment_id, true);
meta_.segmentDescs.emplace_back(segment_desc);
meta_.store = store;
meta_.backendIndex = backendIndex_;
meta_.bufferBaseIndex = backendIndex_ * 10;
while (nextRankForConnection_ != size_) {
std::this_thread::sleep_for(std::chrono::milliseconds(50));
}
// Let the default process group warm up the transfer engine
if (backendIndex_ == 0) {
std::vector<TransferRequest> entries;
for (int i = rank_; i < size_; ++i) {
entries.push_back(TransferRequest{
.opcode = TransferRequest::READ,
.source =
(int32_t*)meta_.segmentDescs[rank_]->buffers[4].addr + 1,
.target_id = meta_.segmentIDs[i],
.target_offset = meta_.segmentDescs[i]->buffers[6].addr,
.length = sizeof(int32_t),
});
}
store->set("backend_warmup_" + std::to_string(backendIndex_) + "_" +
std::to_string(rank_),
"1");
// Ensure all backends have received peer data
for (int i = 0; i < size_; i++) {
store->get_to_str("backend_warmup_" +
std::to_string(backendIndex_) + "_" +
std::to_string(i));
}
auto batchID = engine_.allocateBatchID(entries.size());
engine_.submitTransfer(batchID, entries);
if (options && options->isExtension_) {
auto key = "extension_task_count_" + std::to_string(backendIndex_) +
"_" + std::to_string(rank_);
while (true) {
bool batch_done = true;
TransferStatus status;
for (int i = 0; i < size_ - rank_; ++i) {
engine_.getTransferStatus(batchID, i, status);
if (status.s != TransferStatusEnum::COMPLETED &&
status.s != TransferStatusEnum::FAILED) {
batch_done = false;
break;
}
}
if (batch_done) {
if (store->check({key})) {
meta_.taskCount = std::atoi((char*)store->get(key).data());
break;
}
std::this_thread::sleep_for(std::chrono::milliseconds(50));
}
}
store->set("backend_init_" + std::to_string(backendIndex_) + "_" +
std::to_string(rank_),
"1");
// Ensure that all ranks have been initialized
for (int i = 0; i < size_; i++) {
store->get_to_str("backend_init_" + std::to_string(backendIndex_) +
"_" + std::to_string(i));
}
store->deleteKey("server_name_" + std::to_string(backendIndex_) + "_" +
std::to_string(rank_));
// Increment backend index
++backendIndex_;
@ -256,13 +207,13 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::broadcast(
at::cuda::getCurrentCUDAStream(tensor.device().index());
return worker_.putTaskCuda(
c10d::OpType::BROADCAST, tensorSize, root, &meta_, stream,
[&](void* dst, size_t pos, size_t realSize) {
[=](void* dst, size_t pos, size_t realSize) {
if (isRoot) {
cudaMemcpyAsync(dst, (char*)tensor.data_ptr() + pos,
realSize, cudaMemcpyHostToDevice, stream);
}
},
[&](void* src, size_t pos, size_t realSize) {
[=](void* src, size_t pos, size_t realSize) {
cudaMemcpyAsync((char*)tensor.data_ptr() + pos, src, realSize,
cudaMemcpyDeviceToHost, stream);
});
@ -276,7 +227,7 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::allreduce(
auto tensor = tensors.back();
size_t tensorSize = tensor.numel() * tensor.element_size();
if (isCpu_) {
auto numRanks = size_;
auto numRanks = meta_.size;
return worker_.putTaskCpu(
c10d::OpType::ALLREDUCE, tensorSize, 0, &meta_,
[=](void* dst, size_t pos, size_t realSize) {
@ -285,20 +236,20 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::allreduce(
[=](void* src, size_t pos, size_t realSize) {
memset((char*)tensor.data_ptr() + pos, 0, realSize);
launchReduceCpu(tensor, pos, realSize, src, numRanks,
opts.reduceOp);
opts.reduceOp, meta_.activeRanks);
});
} else {
auto stream = at::cuda::getCurrentCUDAStream(tensor.device().index());
return worker_.putTaskCuda(
c10d::OpType::ALLREDUCE, tensorSize, 0, &meta_, stream,
[&](void* dst, size_t pos, size_t realSize) {
[=](void* dst, size_t pos, size_t realSize) {
cudaMemcpyAsync(dst, (char*)tensor.data_ptr() + pos, realSize,
cudaMemcpyHostToDevice, stream);
},
[&](void* src, size_t pos, size_t realSize) {
[=](void* src, size_t pos, size_t realSize) {
cudaMemsetAsync((char*)tensor.data_ptr() + pos, 0, realSize,
stream);
launchReduceKernel(tensor, pos, realSize, src, size_,
launchReduceKernel(tensor, pos, realSize, src, meta_.size,
opts.reduceOp, meta_.activeRanksDevice,
stream);
});
@ -330,11 +281,11 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::allgather(
at::cuda::getCurrentCUDAStream(inputTensor.device().index());
return worker_.putTaskCuda(
c10d::OpType::ALLGATHER, tensorSize, 0, &meta_, stream,
[&](void* dst, size_t pos, size_t realSize) {
[=](void* dst, size_t pos, size_t realSize) {
cudaMemcpyAsync(dst, (char*)inputTensor.data_ptr() + pos,
realSize, cudaMemcpyHostToDevice, stream);
},
[&](void* src, size_t pos, size_t realSize) {
[=](void* src, size_t pos, size_t realSize) {
for (const auto j : c10::irange(outputTensors_.size())) {
cudaMemcpyAsync((char*)outputTensors_[j].data_ptr() + pos,
(char*)src + j * realSize, realSize,
@ -349,7 +300,7 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::_allgather_base(
const c10d::AllgatherOptions& opts) {
size_t tensorSize = inputBuffer.numel() * inputBuffer.element_size();
if (isCpu_) {
auto numRanks = size_;
auto numRanks = meta_.size;
return worker_.putTaskCpu(
c10d::OpType::_ALLGATHER_BASE, tensorSize, 0, &meta_,
[=](void* dst, size_t pos, size_t realSize) {
@ -367,12 +318,12 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::_allgather_base(
at::cuda::getCurrentCUDAStream(inputBuffer.device().index());
return worker_.putTaskCuda(
c10d::OpType::_ALLGATHER_BASE, tensorSize, 0, &meta_, stream,
[&](void* dst, size_t pos, size_t realSize) {
[=](void* dst, size_t pos, size_t realSize) {
cudaMemcpyAsync(dst, (char*)inputBuffer.data_ptr() + pos,
realSize, cudaMemcpyHostToDevice, stream);
},
[&](void* src, size_t pos, size_t realSize) {
for (const auto j : c10::irange(size_)) {
[=](void* src, size_t pos, size_t realSize) {
for (const auto j : c10::irange(meta_.size)) {
cudaMemcpyAsync(
(char*)outputBuffer.data_ptr() + j * tensorSize + pos,
(char*)src + j * realSize, realSize,
@ -387,7 +338,7 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::_reduce_scatter_base(
const c10d::ReduceScatterOptions& opts) {
size_t tensorSize = outputBuffer.numel() * outputBuffer.element_size();
if (isCpu_) {
auto numRanks = size_;
auto numRanks = meta_.size;
return worker_.putTaskCpu(
c10d::OpType::_REDUCE_SCATTER_BASE, tensorSize, 0, &meta_,
[=](void* dst, size_t pos, size_t realSize) {
@ -400,25 +351,25 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::_reduce_scatter_base(
[=](void* src, size_t pos, size_t realSize) {
memset((char*)outputBuffer.data_ptr() + pos, 0, realSize);
launchReduceCpu(outputBuffer, pos, realSize, src, numRanks,
opts.reduceOp);
opts.reduceOp, meta_.activeRanks);
});
} else {
auto stream =
at::cuda::getCurrentCUDAStream(inputBuffer.device().index());
return worker_.putTaskCuda(
c10d::OpType::_REDUCE_SCATTER_BASE, tensorSize, 0, &meta_, stream,
[&](void* dst, size_t pos, size_t realSize) {
for (const auto j : c10::irange(size_)) {
[=](void* dst, size_t pos, size_t realSize) {
for (const auto j : c10::irange(meta_.size)) {
cudaMemcpyAsync(
(char*)dst + j * realSize,
(char*)inputBuffer.data_ptr() + j * tensorSize + pos,
realSize, cudaMemcpyHostToDevice, stream);
}
},
[&](void* src, size_t pos, size_t realSize) {
[=](void* src, size_t pos, size_t realSize) {
cudaMemsetAsync((char*)outputBuffer.data_ptr() + pos, 0,
realSize, stream);
launchReduceKernel(outputBuffer, pos, realSize, src, size_,
launchReduceKernel(outputBuffer, pos, realSize, src, meta_.size,
opts.reduceOp, meta_.activeRanksDevice,
stream);
});
@ -450,14 +401,14 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::alltoall(
at::cuda::getCurrentCUDAStream(inputTensors[0].device().index());
return worker_.putTaskCuda(
c10d::OpType::ALLTOALL, tensorSize, 0, &meta_, stream,
[&](void* dst, size_t pos, size_t realSize) {
[=](void* dst, size_t pos, size_t realSize) {
for (const auto j : c10::irange(inputTensors.size())) {
cudaMemcpyAsync(dst + j * realSize,
(char*)inputTensors[j].data_ptr() + pos,
realSize, cudaMemcpyHostToDevice, stream);
}
},
[&](void* src, size_t pos, size_t realSize) {
[=](void* src, size_t pos, size_t realSize) {
for (const auto j : c10::irange(outputTensors.size())) {
cudaMemcpyAsync((char*)outputTensors[j].data_ptr() + pos,
(char*)src + j * realSize, realSize,
@ -477,6 +428,11 @@ c10::intrusive_ptr<c10d::Work> MooncakeBackend::barrier(
}
void MooncakeBackend::shutdown() {
isShutdown_ = true;
engine_.unregisterLocalMemory(warmup_send_region_);
engine_.unregisterLocalMemory(warmup_recv_region_);
delete[] warmup_send_region_;
delete[] warmup_recv_region_;
for (size_t i = 0; i < 2; i++) {
engine_.unregisterLocalMemory(cpu_sync_send_region_[i]);
engine_.unregisterLocalMemory(cpu_sync_recv_region_[i]);
@ -492,6 +448,162 @@ void MooncakeBackend::shutdown() {
cudaFree(recv_buffer_[i]);
}
}
--backendIndex_;
}
int MooncakeBackend::getNumSyncedRanks() {
std::vector<at::Tensor> tensors;
tensors.emplace_back(torch::tensor(
nextRankForConnection_,
torch::dtype(torch::kInt).device(isCpu_ ? torch::kCPU : torch::kCUDA)));
c10d::AllreduceOptions opts{
.reduceOp = c10d::ReduceOp::MIN,
};
auto work = allreduce(tensors, opts);
work->wait();
if (!isCpu_) {
auto stream =
at::cuda::getCurrentCUDAStream(tensors[0].device().index());
cudaStreamSynchronize(stream);
}
return tensors[0].cpu().item<int>();
}
void MooncakeBackend::extendGroupSizeTo(int size) {
LOG(INFO) << rank_ << " extend to " << size;
meta_.size = size;
meta_.taskCount = 0;
// TODO: compatibility with fault-tolerance
meta_.activeRanksTensor =
at::ones({size}, torch::dtype(torch::kInt32)
.device(isCpu_ ? torch::kCPU : torch::kCUDA));
}
std::vector<bool> MooncakeBackend::getPeerState(const std::vector<int>& ranks) {
bool activeRanksBackup[kMaxNumRanks];
while (true) {
std::vector<int> input;
for (const int rank : ranks) {
input.push_back(meta_.peerConnected[rank]);
}
for (int i = 0; i < meta_.size; i++) {
activeRanksBackup[i] = meta_.activeRanks[i];
}
std::vector<at::Tensor> tensors;
tensors.emplace_back(torch::tensor(
input, torch::dtype(torch::kInt)
.device(isCpu_ ? torch::kCPU : torch::kCUDA)));
c10d::AllreduceOptions opts{
.reduceOp = c10d::ReduceOp::MIN,
};
auto work = allreduce(tensors, opts);
work->wait();
if (!isCpu_) {
auto stream =
at::cuda::getCurrentCUDAStream(tensors[0].device().index());
cudaStreamSynchronize(stream);
}
bool activeRanksChanged = false;
for (int i = 0; i < meta_.size; i++) {
if (activeRanksBackup[i] != meta_.activeRanks[i]) {
activeRanksChanged = true;
break;
}
}
if (!activeRanksChanged) {
std::vector<bool> output;
for (int i = 0; i < tensors[0].size(0); ++i) {
output.push_back(tensors[0].cpu()[i].item<int>() != 0);
}
return output;
}
}
}
void MooncakeBackend::recoverRanks(const std::vector<int>& ranks) {
for (const int rank : ranks) {
TORCH_CHECK(meta_.peerConnected[rank]);
meta_.activeRanks[rank] = true;
meta_.store->set("extension_task_count_" +
std::to_string(meta_.backendIndex) + "_" +
std::to_string(rank),
std::to_string(meta_.taskCount));
}
}
void MooncakeBackend::connectionPoller(c10::intrusive_ptr<::c10d::Store> store,
int backendIndex) {
while (!isShutdown_) {
for (int pollingRank = 0; pollingRank <= nextRankForConnection_;
++pollingRank) {
if (meta_.peerConnected[pollingRank]) {
continue;
}
std::string serverNameKey = "server_name_" +
std::to_string(backendIndex) + "_" +
std::to_string(pollingRank);
try {
if (!store->check({serverNameKey})) {
continue;
}
} catch (const std::exception& e) {
std::this_thread::sleep_for(std::chrono::milliseconds(50));
continue;
}
if (isShutdown_) {
break;
}
auto peerServerName = store->get_to_str(serverNameKey);
auto segment_id = engine_.openSegment(peerServerName);
meta_.segmentIDs[pollingRank] = segment_id;
auto segment_desc =
engine_.getMetadata()->getSegmentDescByID(segment_id, true);
meta_.segmentDescs[pollingRank] = segment_desc;
if (backendIndex == 0) {
if (pollingRank <= rank_) {
// Send a pre-flight request to establish connections
std::vector<TransferRequest> entries;
auto batchID = engine_.allocateBatchID(1);
engine_.submitTransfer(
batchID,
{TransferRequest{
.opcode = TransferRequest::WRITE,
.source = warmup_send_region_,
.target_id = meta_.segmentIDs[pollingRank],
.target_offset = meta_.segmentDescs[pollingRank]
->buffers[9]
.addr +
rank_ * sizeof(int32_t),
.length = sizeof(int32_t),
}});
while (true) {
TransferStatus status;
engine_.getTransferStatus(batchID, 0, status);
if (status.s == TransferStatusEnum::COMPLETED) {
break;
} else if (status.s == TransferStatusEnum::FAILED) {
LOG(WARNING) << "Warmup request " << rank_ << " -> "
<< pollingRank << " failed.";
break;
}
}
} else {
// Wait for the warmup signals
while (!warmup_recv_region_[pollingRank]) {
std::this_thread::sleep_for(
std::chrono::milliseconds(50));
}
}
}
meta_.peerConnected[pollingRank] = true;
if (pollingRank == nextRankForConnection_) {
++nextRankForConnection_;
}
}
std::this_thread::sleep_for(std::chrono::milliseconds(50));
}
}
} // namespace mooncake

View File

@ -63,65 +63,85 @@ __global__ void enqueueTaskKernel(c10d::OpType opType, size_t tensorSize,
template <typename scalar_t>
__global__ void reduceKernel(scalar_t* dst, const scalar_t* src,
size_t numElements, size_t numRanks,
bool* activeRanks) {
c10d::ReduceOp::RedOpType op, bool* activeRanks) {
size_t thread_idx = blockIdx.x * blockDim.x + threadIdx.x;
size_t stride = blockDim.x * gridDim.x;
for (size_t elem_idx = thread_idx; elem_idx < numElements;
elem_idx += stride) {
scalar_t sum = 0;
bool valid = false;
scalar_t acc = 0;
for (size_t rank = 0; rank < numRanks; ++rank) {
if (activeRanks[rank]) {
sum += src[rank * numElements + elem_idx];
if (!valid) {
acc = src[elem_idx];
valid = true;
} else {
switch (op) {
case c10d::ReduceOp::SUM:
acc += src[rank * numElements + elem_idx];
break;
case c10d::ReduceOp::MIN:
acc = std::min(src[rank * numElements + elem_idx],
acc);
break;
default:
// never
}
}
}
}
dst[elem_idx] = sum;
dst[elem_idx] = acc;
}
}
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.");
TORCH_CHECK(op == c10d::ReduceOp::SUM || op == c10d::ReduceOp::MIN,
"Only support SUM/MIN 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>>>((uint8_t*)ptr, (uint8_t*)src,
num, numRanks, activeRanks);
num, numRanks, op.op_,
activeRanks);
break;
case c10::kChar:
reduceKernel<<<64, 256, 0, stream>>>((int8_t*)ptr, (int8_t*)src,
num, numRanks, activeRanks);
reduceKernel<<<64, 256, 0, stream>>>(
(int8_t*)ptr, (int8_t*)src, num, numRanks, op.op_, activeRanks);
break;
case c10::kShort:
reduceKernel<<<64, 256, 0, stream>>>((int16_t*)ptr, (int16_t*)src,
num, numRanks, activeRanks);
num, numRanks, op.op_,
activeRanks);
break;
case c10::kInt:
reduceKernel<<<64, 256, 0, stream>>>((int*)ptr, (int*)src, num,
numRanks, activeRanks);
numRanks, op.op_, activeRanks);
break;
case c10::kLong:
reduceKernel<<<64, 256, 0, stream>>>((int64_t*)ptr, (int64_t*)src,
num, numRanks, activeRanks);
num, numRanks, op.op_,
activeRanks);
break;
case c10::kFloat:
reduceKernel<<<64, 256, 0, stream>>>((float*)ptr, (float*)src, num,
numRanks, activeRanks);
numRanks, op.op_, activeRanks);
break;
case c10::kDouble:
reduceKernel<<<64, 256, 0, stream>>>((double*)ptr, (double*)src,
num, numRanks, activeRanks);
reduceKernel<<<64, 256, 0, stream>>>(
(double*)ptr, (double*)src, num, numRanks, op.op_, activeRanks);
break;
case c10::kBool:
reduceKernel<<<64, 256, 0, stream>>>((bool*)ptr, (bool*)src, num,
numRanks, activeRanks);
numRanks, op.op_, activeRanks);
break;
case c10::kBFloat16:
reduceKernel<<<64, 256, 0, stream>>>((at::BFloat16*)ptr,
(at::BFloat16*)src, num,
numRanks, activeRanks);
numRanks, op.op_, activeRanks);
break;
default:
TORCH_CHECK(false, c10::str("Unsupported reduce dtype: ",
@ -147,12 +167,21 @@ T applyReduceOp(const T& a, const T& b, c10d::ReduceOp op) {
template <typename T>
void reduceCpu(T* dst, const T* src, size_t numElements, size_t numRanks,
c10d::ReduceOp op) {
c10d::ReduceOp op, bool* activeRanks) {
at::parallel_for(0, numElements, 1024, [&](int64_t begin, int64_t end) {
for (int64_t i = begin; i < end; ++i) {
bool valid = false;
T acc = src[i];
for (int64_t rank = 1; rank < numRanks; ++rank) {
acc = applyReduceOp(acc, src[i + rank * numElements], op);
for (int64_t rank = 0; rank < numRanks; ++rank) {
if (activeRanks[rank]) {
if (!valid) {
acc = src[i];
valid = true;
} else {
acc =
applyReduceOp(acc, src[i + rank * numElements], op);
}
}
}
dst[i] = acc;
}
@ -160,34 +189,39 @@ void reduceCpu(T* dst, const T* src, size_t numElements, size_t numRanks,
}
void launchReduceCpu(at::Tensor dst, size_t pos, size_t realSize, void* src,
size_t numRanks, c10d::ReduceOp op) {
size_t numRanks, c10d::ReduceOp op, bool* activeRanks) {
auto ptr = (char*)dst.data_ptr() + pos;
size_t num = realSize / dst.element_size();
switch (dst.scalar_type()) {
case c10::kByte:
reduceCpu((uint8_t*)ptr, (uint8_t*)src, num, numRanks, op);
reduceCpu((uint8_t*)ptr, (uint8_t*)src, num, numRanks, op,
activeRanks);
break;
case c10::kChar:
reduceCpu((int8_t*)ptr, (int8_t*)src, num, numRanks, op);
reduceCpu((int8_t*)ptr, (int8_t*)src, num, numRanks, op,
activeRanks);
break;
case c10::kShort:
reduceCpu((int16_t*)ptr, (int16_t*)src, num, numRanks, op);
reduceCpu((int16_t*)ptr, (int16_t*)src, num, numRanks, op,
activeRanks);
break;
case c10::kInt:
reduceCpu((int*)ptr, (int*)src, num, numRanks, op);
reduceCpu((int*)ptr, (int*)src, num, numRanks, op, activeRanks);
break;
case c10::kLong:
reduceCpu((int64_t*)ptr, (int64_t*)src, num, numRanks, op);
reduceCpu((int64_t*)ptr, (int64_t*)src, num, numRanks, op,
activeRanks);
break;
case c10::kFloat:
reduceCpu((float*)ptr, (float*)src, num, numRanks, op);
reduceCpu((float*)ptr, (float*)src, num, numRanks, op, activeRanks);
break;
case c10::kDouble:
reduceCpu((double*)ptr, (double*)src, num, numRanks, op);
reduceCpu((double*)ptr, (double*)src, num, numRanks, op,
activeRanks);
break;
case c10::kBool:
reduceCpu((bool*)ptr, (bool*)src, num, numRanks, op);
reduceCpu((bool*)ptr, (bool*)src, num, numRanks, op, activeRanks);
break;
default:
TORCH_CHECK(false, c10::str("Unsupported reduce dtype: ",

View File

@ -119,7 +119,16 @@ void MooncakeWorker::startWorker() {
<< " marking peer " << j
<< " as broken during transferring op "
<< (int)task.opType;
group->store->deleteKey(
"server_name_" +
std::to_string(group->backendIndex) +
"_" + std::to_string(j));
group->store->deleteKey(
"extension_task_count_" +
std::to_string(group->backendIndex) +
"_" + std::to_string(j));
group->activeRanks[j] = false;
group->peerConnected[j] = false;
} else {
batch_done = false;
break;
@ -177,7 +186,16 @@ void MooncakeWorker::startWorker() {
<< " marking peer " << j
<< " as broken during syncing op "
<< (int)task.opType;
group->store->deleteKey(
"server_name_" +
std::to_string(group->backendIndex) + "_" +
std::to_string(j));
group->store->deleteKey(
"extension_task_count_" +
std::to_string(group->backendIndex) + "_" +
std::to_string(j));
group->activeRanks[j] = false;
group->peerConnected[j] = false;
} else {
all_received = false;
break;

View File

@ -58,6 +58,32 @@ at::Tensor getActiveRanks(c10::intrusive_ptr<c10d::Backend> backend) {
return mooncakeBackend->getActiveRanksTensor();
}
int getNumSyncedRanks(c10::intrusive_ptr<c10d::Backend> backend) {
auto mooncakeBackend =
c10::static_intrusive_pointer_cast<MooncakeBackend>(backend);
return mooncakeBackend->getNumSyncedRanks();
}
void extendGroupSizeTo(c10::intrusive_ptr<c10d::Backend> backend, int size) {
auto mooncakeBackend =
c10::static_intrusive_pointer_cast<MooncakeBackend>(backend);
mooncakeBackend->extendGroupSizeTo(size);
}
std::vector<bool> getPeerState(c10::intrusive_ptr<c10d::Backend> backend,
const std::vector<int> &ranks) {
auto mooncakeBackend =
c10::static_intrusive_pointer_cast<MooncakeBackend>(backend);
return mooncakeBackend->getPeerState(ranks);
}
void recoverRanks(c10::intrusive_ptr<c10d::Backend> backend,
const std::vector<int> &ranks) {
auto mooncakeBackend =
c10::static_intrusive_pointer_cast<MooncakeBackend>(backend);
mooncakeBackend->recoverRanks(ranks);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("createMooncakeBackend", &createMooncakeBackend);
m.def("createMooncakeCpuBackend", &createMooncakeCpuBackend);
@ -65,11 +91,17 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("set_device_filter", &MooncakeBackend::setDeviceFilter);
m.def("get_preferred_hca", &getPreferredHca);
m.def("get_active_ranks", &getActiveRanks);
m.def("get_num_synced_ranks", &getNumSyncedRanks);
m.def("extend_group_size_to", &extendGroupSizeTo);
m.def("get_peer_state", &getPeerState);
m.def("recover_ranks", &recoverRanks);
py::class_<MooncakeBackend::MooncakeBackendOptions,
c10::intrusive_ptr<MooncakeBackend::MooncakeBackendOptions>>(
m, "MooncakeBackendOptions")
.def(py::init<at::Tensor>(), py::arg("active_ranks"));
.def(py::init<at::Tensor>(), py::arg("active_ranks"))
.def(py::init<at::Tensor, bool>(), py::arg("active_ranks"),
py::arg("is_extension"));
m.def("get_ep_buffer_size_hint", &get_ep_buffer_size_hint);

View File

@ -0,0 +1,137 @@
import os
import time
import unittest
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from mooncake import ep
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = "19000"
broken_rank = 1
def _elastic_worker(rank, num_processes, signals):
"""Worker for testing elastic world size extension."""
assert num_processes % 2 == 0
if rank < num_processes // 2:
# Ensure correct operation before extension
world_size = num_processes // 2
dist.init_process_group(
backend="mooncake-cpu",
rank=rank,
world_size=world_size,
)
tensor = torch.tensor([rank + 1], dtype=torch.int32, device="cpu")
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
assert tensor.item() == sum(range(1, world_size + 1))
if rank == 0:
signals["extend"] = 1
backend = dist.group.WORLD._get_backend(torch.device("cpu"))
while True:
num_synced_ranks = ep.get_num_synced_ranks(backend)
if num_synced_ranks == num_processes:
break
# Simulate ongoing operations
tensor = torch.tensor([rank + 1], dtype=torch.int32, device="cpu")
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
assert tensor.item() == sum(range(1, world_size + 1))
# Extend world
ep.extend_group_size_to(backend, num_processes)
else:
while "extend" not in signals:
time.sleep(1)
dist.init_process_group(
backend="mooncake-cpu",
rank=rank,
world_size=num_processes,
)
# Ensure correct operation after extension
tensor = torch.tensor([rank + 1], dtype=torch.int32, device="cpu")
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
assert tensor.item() == sum(range(1, num_processes + 1)), (
f"Rank {rank} expected {sum(range(1, num_processes + 1))}, "
f"get {tensor.item()}"
)
def _recovery_worker(rank, num_processes, signals):
"""Worker for testing rank recovery."""
if rank < num_processes:
dist.init_process_group(
backend="mooncake-cpu",
rank=rank,
world_size=num_processes,
)
if rank == broken_rank:
return # Simulate broken rank
tensor = torch.tensor([rank], dtype=torch.int32, device="cpu")
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
assert tensor.item() == sum(range(0, num_processes)) - broken_rank
time.sleep(5)
signals["recover"] = 1
backend = dist.group.WORLD._get_backend(torch.device("cpu"))
while True:
(peer_state,) = ep.get_peer_state(backend, [broken_rank])
if peer_state:
break
ep.recover_ranks(backend, [broken_rank])
# Ensure correct operation after recovery
tensor = torch.tensor([rank], dtype=torch.int32, device="cpu")
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
assert tensor.item() == sum(range(0, num_processes)), (
f"Rank {rank} expected {sum(range(0, num_processes))}, "
f"get {tensor.item()}"
)
else:
while "recover" not in signals:
time.sleep(1)
dist.init_process_group(
backend="mooncake-cpu",
rank=broken_rank,
world_size=num_processes,
pg_options=ep.MooncakeBackendOptions(
torch.ones((num_processes,), dtype=torch.int32),
True,
),
)
# Ensure correct operation after recovery
tensor = torch.tensor([broken_rank], dtype=torch.int32, device="cpu")
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
assert tensor.item() == sum(range(0, num_processes)), (
f"Rank {rank} expected {sum(range(0, num_processes))}, "
f"get {tensor.item()}"
)
class TestMooncakeBackend(unittest.TestCase):
def test_elastic_extension(self):
num_processes = 4
mp_manager = mp.Manager()
signals = mp_manager.dict()
mp.spawn(_elastic_worker, args=(num_processes, signals), nprocs=num_processes)
def test_rank_recovery(self):
num_processes = 4
mp_manager = mp.Manager()
signals = mp_manager.dict()
mp.spawn(
_recovery_worker,
args=(num_processes, signals),
nprocs=num_processes + 1,
)
if __name__ == "__main__":
unittest.main()