From 8ff33590d2914f44364bf28d22348d5381bace37 Mon Sep 17 00:00:00 2001 From: UNIDY2002 Date: Wed, 25 Mar 2026 15:58:57 +0800 Subject: [PATCH 1/2] Implement rank-recovery deferred join --- mooncake-ep/include/mooncake_ep_buffer.h | 9 +- mooncake-ep/src/mooncake_ep_buffer.cpp | 24 ++- mooncake-pg/include/connection_poller.h | 16 +- mooncake-pg/include/mooncake_backend.h | 9 + mooncake-pg/src/connection_poller.cpp | 45 ++++- mooncake-pg/src/mooncake_backend.cpp | 162 ++++++++++++++---- mooncake-pg/src/mooncake_worker_thread.cpp | 2 + mooncake-pg/src/pg_py.cpp | 7 + mooncake-wheel/mooncake/mooncake_ep_buffer.py | 15 +- .../tests/test_mooncake_backend_elastic.py | 41 ++++- 10 files changed, 278 insertions(+), 52 deletions(-) diff --git a/mooncake-ep/include/mooncake_ep_buffer.h b/mooncake-ep/include/mooncake_ep_buffer.h index 69e0729abd..3cbb5d0beb 100644 --- a/mooncake-ep/include/mooncake_ep_buffer.h +++ b/mooncake-ep/include/mooncake_ep_buffer.h @@ -173,13 +173,15 @@ struct MooncakeEpBuffer { void sync_ib(const std::vector& remote_addrs, const std::vector& remote_keys, const std::vector& remote_qpns, - const std::vector& remote_lids); + const std::vector& remote_lids, + const std::vector& active_ranks_mask); void sync_roce(const std::vector& remote_addrs, const std::vector& remote_keys, const std::vector& remote_qpns, const std::vector& subnet_prefixes, - const std::vector& interface_ids); + const std::vector& interface_ids, + const std::vector& active_ranks_mask); std::tuple get_mr_info() { return {(int64_t)mr->addr, (int32_t)mr->rkey}; @@ -208,7 +210,8 @@ struct MooncakeEpBuffer { std::vector get_ipc_handle(); void sync_nvlink_ipc_handles( - const std::vector>& remote_handles); + const std::vector>& remote_handles, + const std::vector& active_ranks_mask); }; inline size_t get_ep_buffer_size_hint(int num_max_dispatch_tokens_per_rank, diff --git a/mooncake-ep/src/mooncake_ep_buffer.cpp b/mooncake-ep/src/mooncake_ep_buffer.cpp index e9fb240d0d..555d0d10d8 100644 --- a/mooncake-ep/src/mooncake_ep_buffer.cpp +++ b/mooncake-ep/src/mooncake_ep_buffer.cpp @@ -615,8 +615,11 @@ void MooncakeEpBuffer::update_local_qpns() { void MooncakeEpBuffer::sync_ib(const std::vector& remote_addrs, const std::vector& remote_keys, const std::vector& remote_qpns, - const std::vector& remote_lids) { + const std::vector& remote_lids, + const std::vector& active_ranks_mask) { for (int i = 0; i < USE_QP_COUNT; ++i) { + int peer_rank = i * num_ranks / USE_QP_COUNT; + if (active_ranks_mask[peer_rank] == 0) continue; ibv_ah_attr ah_attr = { .dlid = (uint16_t)remote_lids[i], .port_num = 0, @@ -632,6 +635,7 @@ void MooncakeEpBuffer::sync_ib(const std::vector& remote_addrs, } } for (int i = 0; i < num_ranks; ++i) { + if (active_ranks_mask[i] == 0) continue; uint64_t raddr = i == rank ? (uint64_t)mr->addr : (uint64_t)remote_addrs[i]; cudaMemcpy(raddrs + i * sizeof(uint64_t), &raddr, sizeof(uint64_t), @@ -646,13 +650,14 @@ void MooncakeEpBuffer::sync_roce(const std::vector& remote_addrs, const std::vector& remote_keys, const std::vector& remote_qpns, const std::vector& subnet_prefixes, - const std::vector& interface_ids) { + const std::vector& interface_ids, + const std::vector& active_ranks_mask) { for (int i = 0; i < USE_QP_COUNT; ++i) { + int peer_rank = i * num_ranks / USE_QP_COUNT; + if (active_ranks_mask[peer_rank] == 0) continue; ibv_gid remote_gid{}; - remote_gid.global.subnet_prefix = - subnet_prefixes[i * num_ranks / USE_QP_COUNT]; - remote_gid.global.interface_id = - interface_ids[i * num_ranks / USE_QP_COUNT]; + remote_gid.global.subnet_prefix = subnet_prefixes[peer_rank]; + remote_gid.global.interface_id = interface_ids[peer_rank]; ibv_ah_attr ah_attr = {}; ah_attr.is_global = 1; ah_attr.grh.dgid = remote_gid; @@ -672,6 +677,7 @@ void MooncakeEpBuffer::sync_roce(const std::vector& remote_addrs, } } for (int i = 0; i < num_ranks; ++i) { + if (active_ranks_mask[i] == 0) continue; uint64_t raddr = i == rank ? (uint64_t)mr->addr : (uint64_t)remote_addrs[i]; cudaMemcpy(raddrs + i * sizeof(uint64_t), &raddr, sizeof(uint64_t), @@ -701,7 +707,8 @@ std::vector MooncakeEpBuffer::get_ipc_handle() { } void MooncakeEpBuffer::sync_nvlink_ipc_handles( - const std::vector>& remote_handles) { + const std::vector>& remote_handles, + const std::vector& active_ranks_mask) { int device_count = 0; CUDA_CHECK(cudaGetDeviceCount(&device_count)); @@ -714,6 +721,7 @@ void MooncakeEpBuffer::sync_nvlink_ipc_handles( // handle exchange — cuMemSetAccess already granted all devices // read/write access during allocation. for (int i = 0; i < num_ranks; ++i) { + if (active_ranks_mask[i] == 0) continue; nvlink_array[i] = 1; // Each rank's gdr_buffer is directly accessible; the remote // addresses will be exchanged via the RDMA address sync path @@ -730,6 +738,7 @@ void MooncakeEpBuffer::sync_nvlink_ipc_handles( int group_end = std::min(group_start + device_count, num_ranks); for (int dst_rank = group_start; dst_rank < group_end; ++dst_rank) { + if (active_ranks_mask[dst_rank] == 0) continue; if (dst_rank == rank) { ipc_peer_ptrs_host[dst_rank] = gdr_buffer; continue; @@ -789,6 +798,7 @@ void MooncakeEpBuffer::sync_nvlink_ipc_handles( p2p_ipc_all_enabled_ = true; for (int i = 0; i < num_ranks; ++i) { + if (active_ranks_mask[i] == 0) continue; if (nvlink_array[i] == 0 || ipc_peer_ptrs_host[i] == nullptr) { p2p_ipc_all_enabled_ = false; break; diff --git a/mooncake-pg/include/connection_poller.h b/mooncake-pg/include/connection_poller.h index 85f079739f..481ad2f31d 100644 --- a/mooncake-pg/include/connection_poller.h +++ b/mooncake-pg/include/connection_poller.h @@ -56,6 +56,8 @@ class ConnectionContext { std::atomic groupSize_; + bool isDummy_; + // A mark tracking the group size for which all ranks // in [0, establishedGroupSize_) have been successfully // connected at least once (they may disconnect afterwards). @@ -87,7 +89,7 @@ class ConnectionContext { std::condition_variable backend_wakeup_cv_; public: - ConnectionContext(int backendIndex, int rank, int size, + ConnectionContext(int backendIndex, int rank, int size, bool isDummy, uint64_t* local2global_rank_map, std::string location, c10::intrusive_ptr<::c10d::Store> store, std::shared_ptr meta, @@ -133,6 +135,9 @@ class ConnectionContext { */ void waitUntilAllConnected(); + void bootstrapLocalPeer(const std::string& localServerName, + const SegmentInfo& localRankInfo); + /** * @brief Blocks until all newly added ranks in the * extended group are connected. @@ -148,6 +153,8 @@ class ConnectionContext { void shutdown(); + void setDummy(bool isDummy) { isDummy_ = isDummy; } + static std::string getServerNameStoreKey(int backendIndex, int rank) { return "server_name_" + std::to_string(backendIndex) + "_" + std::to_string(rank); @@ -161,6 +168,11 @@ class ConnectionContext { return "extension_task_count_" + std::to_string(backendIndex) + "_" + std::to_string(rank); } + static std::string getExtensionActiveRanksStoreKey(int backendIndex, + int rank) { + return "extension_active_ranks_" + std::to_string(backendIndex) + "_" + + std::to_string(rank); + } private: // For ConnectionManager @@ -194,6 +206,7 @@ class ConnectionPoller { private: ConnectionPoller(); + void ensureThreadStarted(); void pollerLoop(); bool processContext(const std::shared_ptr& ctx); bool processPeer(const std::shared_ptr& ctx, @@ -202,6 +215,7 @@ class ConnectionPoller { std::mutex wakeup_mutex_; std::condition_variable wakeup_cv_; std::thread pollerThread_; + std::atomic pollerThreadStarted_{false}; std::mutex contexts_mutex_; std::atomic contexts_version_{0}; diff --git a/mooncake-pg/include/mooncake_backend.h b/mooncake-pg/include/mooncake_backend.h index 635c16d163..3ad6415666 100644 --- a/mooncake-pg/include/mooncake_backend.h +++ b/mooncake-pg/include/mooncake_backend.h @@ -150,7 +150,14 @@ class MooncakeBackend final : public ::c10d::Backend { void recoverRanks(const std::vector& ranks); + void joinGroup(); + private: + void waitForExtensionState(); + void publishLocalPeerMetadata(); + void setLocalOnlyActiveRanks(); + void syncActiveRanksTensor(); + static TransferEngine* engine_; std::shared_ptr worker_; static bool engineInitialized_; @@ -166,6 +173,7 @@ class MooncakeBackend final : public ::c10d::Backend { std::shared_ptr meta_; bool isShutdown_{false}; uint64_t local2global_rank_map_[kMaxNumRanks]; + std::string localServerName_; // P2P async infrastructure // p2p_proxy_ is created in MooncakeBackend, but can live longer than @@ -181,6 +189,7 @@ class MooncakeBackend final : public ::c10d::Backend { // Similar to p2p_proxy_, connection_ctx_ is created in MooncakeBackend, but // can live longer than MooncakeBackend. std::shared_ptr connection_ctx_; + bool connectionPollerRegistered_{false}; }; } // namespace mooncake diff --git a/mooncake-pg/src/connection_poller.cpp b/mooncake-pg/src/connection_poller.cpp index 28fc517ccb..5c8c94c755 100644 --- a/mooncake-pg/src/connection_poller.cpp +++ b/mooncake-pg/src/connection_poller.cpp @@ -40,6 +40,7 @@ static bool supportFabricMem() { return true; } ConnectionContext::ConnectionContext(int backendIndex, int rank, int size, + bool isDummy, uint64_t* local2global_rank_map, std::string location, c10::intrusive_ptr<::c10d::Store> store, @@ -49,6 +50,7 @@ ConnectionContext::ConnectionContext(int backendIndex, int rank, int size, : backendIndex_(backendIndex), rank_(rank), groupSize_(size), + isDummy_(isDummy), establishedGroupSize_(0), local2global_rank_map_(local2global_rank_map), store_(std::move(store)), @@ -131,6 +133,9 @@ void ConnectionContext::waitUntilAllConnected() { } void ConnectionContext::waitUntilNewRanksConnected() { + if (isDummy_) { + return; + } const int targetGroupSize = groupSize_.load(std::memory_order_acquire); const int established = establishedGroupSize_.load(std::memory_order_acquire); @@ -155,6 +160,32 @@ void ConnectionContext::waitUntilNewRanksConnected() { establishedGroupSize_.store(targetGroupSize, std::memory_order_release); } +void ConnectionContext::bootstrapLocalPeer(const std::string& localServerName, + const SegmentInfo& localRankInfo) { + auto& peerState = peerStates_[rank_]; + if (peerState.state == PeerConnectionState::CONNECTED) { + return; + } + + auto segment_id = engine_->openSegment(localServerName); + meta_->segmentIDs[rank_] = segment_id; + peerState.segmentId = segment_id; + memcpy(&meta_->segmentInfos[rank_], &localRankInfo, sizeof(SegmentInfo)); + + meta_->peerConnected[rank_] = true; + ConnectionPoller::GetInstance() + .global_peerConnected_[local2global_rank_map_[rank_]] = true; + peerState.state = PeerConnectionState::CONNECTED; + + { + std::lock_guard lock(backend_wakeup_mutex_); + totalConnectedPeers_.store(1, std::memory_order_release); + if (isAllPeerConnected()) { + backend_wakeup_cv_.notify_all(); + } + } +} + void ConnectionContext::shutdown() { // Notify backends that may be blocked in waitUntilAllConnected. { @@ -352,6 +383,8 @@ bool ConnectionContext::pollPeer(int pollingRank) { // reports a failure. We must set both to false here. global_peerConnected_[globalPollingRank] = false; meta_->peerConnected[pollingRank] = false; + meta_->activeRanks[pollingRank] = false; + meta_->activeRanksTensor[pollingRank] = 0; // Reset store store_->deleteKey( @@ -359,6 +392,8 @@ bool ConnectionContext::pollPeer(int pollingRank) { store_->deleteKey(getBufferStoreKey(backendIndex_, pollingRank)); store_->deleteKey( getExtensionTaskCountStoreKey(backendIndex_, pollingRank)); + store_->deleteKey( + getExtensionActiveRanksStoreKey(backendIndex_, pollingRank)); // Reset warmup region *reinterpret_cast( @@ -406,12 +441,20 @@ bool ConnectionContext::tryStop() { return stopped; } -ConnectionPoller::ConnectionPoller() { +ConnectionPoller::ConnectionPoller() = default; + +void ConnectionPoller::ensureThreadStarted() { + bool expected = false; + if (!pollerThreadStarted_.compare_exchange_strong( + expected, true, std::memory_order_acq_rel)) { + return; + } pollerThread_ = std::thread([this] { pollerLoop(); }); } void ConnectionPoller::registerContext( const std::shared_ptr& ctx) { + ensureThreadStarted(); { std::lock_guard lock(contexts_mutex_); contexts_.push_back(ctx); diff --git a/mooncake-pg/src/mooncake_backend.cpp b/mooncake-pg/src/mooncake_backend.cpp index 35c7c04cca..db3fc45986 100644 --- a/mooncake-pg/src/mooncake_backend.cpp +++ b/mooncake-pg/src/mooncake_backend.cpp @@ -31,6 +31,27 @@ TransferEngine* MooncakeBackend::engine_ = new TransferEngine(true); bool MooncakeBackend::engineInitialized_ = false; int MooncakeBackend::backendIndex_ = 0; +namespace { + +std::vector serializeActiveRanks(const bool* activeRanks, int size) { + std::vector bytes(size); + for (int i = 0; i < size; ++i) { + bytes[i] = activeRanks[i] ? 1 : 0; + } + return bytes; +} + +void deserializeActiveRanks(const std::vector& bytes, + bool* activeRanks, int size) { + TORCH_CHECK(static_cast(bytes.size()) == size, + "Unexpected active-ranks snapshot size."); + for (int i = 0; i < size; ++i) { + activeRanks[i] = (bytes[i] != 0); + } +} + +} // namespace + // Async Work implementation for P2P operations processed by worker threads. class MooncakeP2PWork : public ::c10d::Work { public: @@ -98,7 +119,7 @@ MooncakeBackend::MooncakeBackend( engine_->init(P2PHANDSHAKE, hostIp_); engineInitialized_ = true; } - std::string localServerName = engine_->getLocalIpAndPort(); + localServerName_ = engine_->getLocalIpAndPort(); // construct local to global rank map if (globalRanks.size() == static_cast(size)) { for (int i = 0; i < size; ++i) { @@ -196,8 +217,8 @@ MooncakeBackend::MooncakeBackend( meta_ = std::make_shared(); connection_ctx_ = std::make_shared( - backendIndex_, rank, size, local2global_rank_map_, location, store, - meta_, p2p_proxy_, engine_); + backendIndex_, rank, size, options_ && options_->isExtension_, + local2global_rank_map_, location, store, meta_, p2p_proxy_, engine_); rank_info.send_buffer[0] = (uint64_t)send_buffer_[0]; rank_info.send_buffer[1] = (uint64_t)send_buffer_[1]; @@ -219,16 +240,6 @@ MooncakeBackend::MooncakeBackend( // Sync metadata std::vector rank_info_bytes(sizeof(SegmentInfo)); memcpy(rank_info_bytes.data(), &rank_info, sizeof(SegmentInfo)); - auto bufferKey = ConnectionContext::getBufferStoreKey(backendIndex_, rank_); - store->set(bufferKey, rank_info_bytes); - - auto serverNameKey = - ConnectionContext::getServerNameStoreKey(backendIndex_, rank_); - store->set(serverNameKey, localServerName); - - // Start polling connection - ConnectionPoller::GetInstance().registerContext(connection_ctx_); - meta_->rank = rank; meta_->size = size; meta_->taskCount = 0; @@ -265,21 +276,14 @@ MooncakeBackend::MooncakeBackend( meta_->bufferBaseIndex = backendIndex_ * 10; p2p_proxy_->BindMeta(meta_); - // Wait for peers - connection_ctx_->waitUntilAllConnected(); - + connection_ctx_->bootstrapLocalPeer(localServerName_, rank_info); if (options_ && options_->isExtension_) { - auto key = ConnectionContext::getExtensionTaskCountStoreKey( - backendIndex_, rank_); - while (true) { - if (store->check({key})) { - auto data = store->get(key); - std::string val(data.begin(), data.end()); - meta_->taskCount = std::stoi(val); - break; - } - std::this_thread::sleep_for(std::chrono::milliseconds(50)); - } + setLocalOnlyActiveRanks(); + } else { + publishLocalPeerMetadata(); + ConnectionPoller::GetInstance().registerContext(connection_ctx_); + connectionPollerRegistered_ = true; + connection_ctx_->waitUntilAllConnected(); } // Increment backend index @@ -761,7 +765,10 @@ void MooncakeBackend::shutdown() { p2p_proxy_.reset(); connection_ctx_->shutdown(); - ConnectionPoller::GetInstance().removeContext(connection_ctx_); + if (connectionPollerRegistered_) { + ConnectionPoller::GetInstance().removeContext(connection_ctx_); + connectionPollerRegistered_ = false; + } for (size_t i = 0; i < 2; i++) { engine_->unregisterLocalMemory(cpu_sync_send_region_[i]); @@ -780,6 +787,76 @@ void MooncakeBackend::shutdown() { } } +void MooncakeBackend::syncActiveRanksTensor() { + std::vector active_ranks(meta_->size); + for (int i = 0; i < meta_->size; ++i) { + active_ranks[i] = meta_->activeRanks[i] ? 1 : 0; + } + + auto cpu_tensor = torch::tensor(active_ranks, torch::dtype(torch::kInt32)); + if (!meta_->activeRanksTensor.defined() || + meta_->activeRanksTensor.size(0) != meta_->size) { + meta_->activeRanksTensor = + cpu_tensor.to(isCpu_ ? torch::kCPU : torch::kCUDA); + return; + } + + if (meta_->activeRanksTensor.device().is_cpu()) { + meta_->activeRanksTensor.copy_(cpu_tensor); + } else { + meta_->activeRanksTensor.copy_( + cpu_tensor.to(meta_->activeRanksTensor.device())); + } +} + +void MooncakeBackend::publishLocalPeerMetadata() { + TORCH_CHECK(meta_->store, + "Publishing local peer metadata requires a valid Store."); + + std::vector rank_info_bytes(sizeof(SegmentInfo)); + memcpy(rank_info_bytes.data(), &rank_info, sizeof(SegmentInfo)); + + auto bufferKey = + ConnectionContext::getBufferStoreKey(meta_->backendIndex, rank_); + meta_->store->set(bufferKey, rank_info_bytes); + + auto serverNameKey = + ConnectionContext::getServerNameStoreKey(meta_->backendIndex, rank_); + meta_->store->set(serverNameKey, localServerName_); +} + +void MooncakeBackend::setLocalOnlyActiveRanks() { + for (int i = 0; i < meta_->size; ++i) { + meta_->activeRanks[i] = (i == meta_->rank); + } + syncActiveRanksTensor(); +} + +void MooncakeBackend::waitForExtensionState() { + TORCH_CHECK(meta_->store, "Recovery join requires a valid Store."); + + auto task_count_key = ConnectionContext::getExtensionTaskCountStoreKey( + meta_->backendIndex, rank_); + auto active_ranks_key = ConnectionContext::getExtensionActiveRanksStoreKey( + meta_->backendIndex, rank_); + + while (true) { + if (meta_->store->check({task_count_key, active_ranks_key})) { + auto task_count_data = meta_->store->get(task_count_key); + std::string task_count(task_count_data.begin(), + task_count_data.end()); + meta_->taskCount = std::stoi(task_count); + + auto active_ranks = meta_->store->get(active_ranks_key); + deserializeActiveRanks(active_ranks, meta_->activeRanks, + meta_->size); + syncActiveRanksTensor(); + return; + } + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + } +} + int MooncakeBackend::getNumSyncedRanks() { std::vector tensors; tensors.emplace_back(torch::tensor( @@ -874,15 +951,38 @@ std::vector MooncakeBackend::getPeerState(const std::vector& ranks) { } void MooncakeBackend::recoverRanks(const std::vector& ranks) { + TORCH_CHECK(meta_->store, "Rank recovery requires a valid Store."); + for (const int rank : ranks) { TORCH_CHECK(rank >= 0 && static_cast(rank) < kMaxNumRanks, "Rank out of range"); TORCH_CHECK(meta_->peerConnected[rank]); meta_->activeRanks[rank] = true; - meta_->store->set("extension_task_count_" + - std::to_string(meta_->backendIndex) + "_" + - std::to_string(rank), + } + + syncActiveRanksTensor(); + auto active_ranks_snapshot = + serializeActiveRanks(meta_->activeRanks, meta_->size); + for (const int rank : ranks) { + meta_->store->set(ConnectionContext::getExtensionTaskCountStoreKey( + meta_->backendIndex, rank), std::to_string(meta_->taskCount)); + meta_->store->set(ConnectionContext::getExtensionActiveRanksStoreKey( + meta_->backendIndex, rank), + active_ranks_snapshot); + } +} + +void MooncakeBackend::joinGroup() { + TORCH_CHECK(options_ && options_->isExtension_, + "joinGroup is only valid for extension backends."); + connection_ctx_->setDummy(false); + publishLocalPeerMetadata(); + if (!connectionPollerRegistered_) { + ConnectionPoller::GetInstance().registerContext(connection_ctx_); + connectionPollerRegistered_ = true; } + connection_ctx_->waitUntilAllConnected(); + waitForExtensionState(); } } // namespace mooncake diff --git a/mooncake-pg/src/mooncake_worker_thread.cpp b/mooncake-pg/src/mooncake_worker_thread.cpp index 29c48d7ca7..002f710a1c 100644 --- a/mooncake-pg/src/mooncake_worker_thread.cpp +++ b/mooncake-pg/src/mooncake_worker_thread.cpp @@ -149,6 +149,7 @@ void MooncakeWorker::startWorker() { // connection poller to reconnect it. group->peerConnected[j] = false; group->activeRanks[j] = false; + group->activeRanksTensor[j] = 0; } else { batch_done = false; break; @@ -230,6 +231,7 @@ void MooncakeWorker::startWorker() { // connection poller to reconnect it. group->peerConnected[j] = false; group->activeRanks[j] = false; + group->activeRanksTensor[j] = 0; } else { task_done = false; break; diff --git a/mooncake-pg/src/pg_py.cpp b/mooncake-pg/src/pg_py.cpp index 74f9b7e72f..78e3269aee 100644 --- a/mooncake-pg/src/pg_py.cpp +++ b/mooncake-pg/src/pg_py.cpp @@ -78,6 +78,12 @@ void recoverRanks(c10::intrusive_ptr backend, mooncakeBackend->recoverRanks(ranks); } +void joinGroup(c10::intrusive_ptr backend) { + auto mooncakeBackend = + c10::static_intrusive_pointer_cast(backend); + mooncakeBackend->joinGroup(); +} + PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("createMooncakeBackend", &createMooncakeBackend); m.def("createMooncakeCpuBackend", &createMooncakeCpuBackend); @@ -89,6 +95,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("extend_group_size_to", &extendGroupSizeTo); m.def("get_peer_state", &getPeerState); m.def("recover_ranks", &recoverRanks); + m.def("join_group", &joinGroup); py::class_>( diff --git a/mooncake-wheel/mooncake/mooncake_ep_buffer.py b/mooncake-wheel/mooncake/mooncake_ep_buffer.py index 2a136475cd..c37502b58e 100644 --- a/mooncake-wheel/mooncake/mooncake_ep_buffer.py +++ b/mooncake-wheel/mooncake/mooncake_ep_buffer.py @@ -150,8 +150,11 @@ def connect(self, is_update: bool = False): dist.all_gather(interface_ids, interface_id, self.group) interface_ids = torch.cat(interface_ids).tolist() + from mooncake.ep import get_active_ranks + active_ranks_mask = get_active_ranks(self.backend).tolist() self.runtime.sync_roce( - raddrs, rkeys, remote_qpns, subnet_prefixes, interface_ids + raddrs, rkeys, remote_qpns, subnet_prefixes, interface_ids, + active_ranks_mask ) else: local_lids = self.runtime.get_local_lids() @@ -169,7 +172,10 @@ def connect(self, is_update: bool = False): dist.all_to_all(remote_lids, local_lids, self.group) remote_lids = torch.cat(remote_lids).tolist() - self.runtime.sync_ib(raddrs, rkeys, remote_qpns, remote_lids) + from mooncake.ep import get_active_ranks + active_ranks_mask = get_active_ranks(self.backend).tolist() + self.runtime.sync_ib(raddrs, rkeys, remote_qpns, remote_lids, + active_ranks_mask) try: local_handle_ints = self.runtime.get_ipc_handle() @@ -183,7 +189,10 @@ def connect(self, is_update: bool = False): ] dist.all_gather(handles, local_handle_tensor, self.group) remote_handles = [h.tolist() for h in handles] - self.runtime.sync_nvlink_ipc_handles(remote_handles) + from mooncake.ep import get_active_ranks + active_ranks_mask = get_active_ranks(self.backend).tolist() + self.runtime.sync_nvlink_ipc_handles(remote_handles, + active_ranks_mask) except Exception as e: import warnings diff --git a/mooncake-wheel/tests/test_mooncake_backend_elastic.py b/mooncake-wheel/tests/test_mooncake_backend_elastic.py index dc9461ac48..8400b2a149 100644 --- a/mooncake-wheel/tests/test_mooncake_backend_elastic.py +++ b/mooncake-wheel/tests/test_mooncake_backend_elastic.py @@ -57,8 +57,8 @@ def _elastic_worker(rank, num_processes, signals): ) -def _recovery_worker(rank, num_processes, signals): - """Worker for testing rank recovery.""" +def _deferred_recovery_worker(rank, num_processes, signals): + """Worker for testing deferred rank recovery join.""" if rank < num_processes: dist.init_process_group( backend="mooncake-cpu", @@ -68,9 +68,10 @@ def _recovery_worker(rank, num_processes, signals): if rank == broken_rank: return # Simulate broken rank + expected_without_broken = sum(range(0, num_processes)) - 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 + assert tensor.item() == expected_without_broken time.sleep(5) signals["recover"] = 1 @@ -79,9 +80,17 @@ def _recovery_worker(rank, num_processes, signals): (peer_state,) = pg.get_peer_state(backend, [broken_rank]) if peer_state: break + + # Healthy ranks keep making progress before the recovered rank joins. + tensor = torch.tensor([rank], dtype=torch.int32, device="cpu") + dist.all_reduce(tensor, op=dist.ReduceOp.SUM) + assert tensor.item() == expected_without_broken + + while "join_ready" not in signals: + time.sleep(0.1) + pg.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)), ( @@ -90,7 +99,8 @@ def _recovery_worker(rank, num_processes, signals): ) else: while "recover" not in signals: - time.sleep(1) + time.sleep(0.1) + dist.init_process_group( backend="mooncake-cpu", rank=broken_rank, @@ -101,7 +111,16 @@ def _recovery_worker(rank, num_processes, signals): ), ) - # Ensure correct operation after recovery + backend = dist.group.WORLD._get_backend(torch.device("cpu")) + + # Deferred join starts in a local-only mode so collectives stay self-contained. + tensor = torch.tensor([broken_rank], dtype=torch.int32, device="cpu") + dist.all_reduce(tensor, op=dist.ReduceOp.SUM) + assert tensor.item() == broken_rank + + signals["join_ready"] = 1 + pg.join_group(backend) + 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)), ( @@ -127,6 +146,16 @@ def test_rank_recovery(self): nprocs=num_processes + 1, ) + def test_rank_recovery_deferred_join(self): + num_processes = 4 + mp_manager = mp.Manager() + signals = mp_manager.dict() + mp.spawn( + _deferred_recovery_worker, + args=(num_processes, signals), + nprocs=num_processes + 1, + ) + if __name__ == "__main__": unittest.main() From c4fb7c0d8e88193f7727acc273f32aae529afc55 Mon Sep 17 00:00:00 2001 From: Xun Sun Date: Sun, 29 Mar 2026 11:35:13 +0800 Subject: [PATCH 2/2] Add memory fraction static parameter to test --- scripts/tone_tests/python/test_moe_mooncake.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/scripts/tone_tests/python/test_moe_mooncake.py b/scripts/tone_tests/python/test_moe_mooncake.py index f6b06a4c33..cce365dfeb 100644 --- a/scripts/tone_tests/python/test_moe_mooncake.py +++ b/scripts/tone_tests/python/test_moe_mooncake.py @@ -39,6 +39,8 @@ def setUpClass(cls): "mooncake", "--mooncake-ib-device", ib_devices, + "--mem-fraction-static", + "0.8", ], )