Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions mooncake-ep/include/mooncake_ep_buffer.h
Original file line number Diff line number Diff line change
Expand Up @@ -173,13 +173,15 @@ struct MooncakeEpBuffer {
void sync_ib(const std::vector<int64_t>& remote_addrs,
const std::vector<int32_t>& remote_keys,
const std::vector<int32_t>& remote_qpns,
const std::vector<int32_t>& remote_lids);
const std::vector<int32_t>& remote_lids,
const std::vector<int>& active_ranks_mask);

void sync_roce(const std::vector<int64_t>& remote_addrs,
const std::vector<int32_t>& remote_keys,
const std::vector<int32_t>& remote_qpns,
const std::vector<int64_t>& subnet_prefixes,
const std::vector<int64_t>& interface_ids);
const std::vector<int64_t>& interface_ids,
const std::vector<int>& active_ranks_mask);

std::tuple<int64_t, int32_t> get_mr_info() {
return {(int64_t)mr->addr, (int32_t)mr->rkey};
Expand Down Expand Up @@ -208,7 +210,8 @@ struct MooncakeEpBuffer {

std::vector<int32_t> get_ipc_handle();
void sync_nvlink_ipc_handles(
const std::vector<std::vector<int32_t>>& remote_handles);
const std::vector<std::vector<int32_t>>& remote_handles,
const std::vector<int>& active_ranks_mask);
};

inline size_t get_ep_buffer_size_hint(int num_max_dispatch_tokens_per_rank,
Expand Down
24 changes: 17 additions & 7 deletions mooncake-ep/src/mooncake_ep_buffer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -615,8 +615,11 @@ void MooncakeEpBuffer::update_local_qpns() {
void MooncakeEpBuffer::sync_ib(const std::vector<int64_t>& remote_addrs,
const std::vector<int32_t>& remote_keys,
const std::vector<int32_t>& remote_qpns,
const std::vector<int32_t>& remote_lids) {
const std::vector<int32_t>& remote_lids,
const std::vector<int>& 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,
Expand All @@ -632,6 +635,7 @@ void MooncakeEpBuffer::sync_ib(const std::vector<int64_t>& 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),
Expand All @@ -646,13 +650,14 @@ void MooncakeEpBuffer::sync_roce(const std::vector<int64_t>& remote_addrs,
const std::vector<int32_t>& remote_keys,
const std::vector<int32_t>& remote_qpns,
const std::vector<int64_t>& subnet_prefixes,
const std::vector<int64_t>& interface_ids) {
const std::vector<int64_t>& interface_ids,
const std::vector<int>& 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;
Expand All @@ -672,6 +677,7 @@ void MooncakeEpBuffer::sync_roce(const std::vector<int64_t>& 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),
Expand Down Expand Up @@ -701,7 +707,8 @@ std::vector<int32_t> MooncakeEpBuffer::get_ipc_handle() {
}

void MooncakeEpBuffer::sync_nvlink_ipc_handles(
const std::vector<std::vector<int32_t>>& remote_handles) {
const std::vector<std::vector<int32_t>>& remote_handles,
const std::vector<int>& active_ranks_mask) {
int device_count = 0;
CUDA_CHECK(cudaGetDeviceCount(&device_count));

Expand All @@ -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
Expand All @@ -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;
Expand Down Expand Up @@ -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;
Expand Down
16 changes: 15 additions & 1 deletion mooncake-pg/include/connection_poller.h
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,8 @@ class ConnectionContext {

std::atomic<int> 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).
Expand Down Expand Up @@ -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<TransferGroupMeta> meta,
Expand Down Expand Up @@ -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.
Expand All @@ -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);
Expand All @@ -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
Expand Down Expand Up @@ -194,6 +206,7 @@ class ConnectionPoller {

private:
ConnectionPoller();
void ensureThreadStarted();
void pollerLoop();
bool processContext(const std::shared_ptr<ConnectionContext>& ctx);
bool processPeer(const std::shared_ptr<ConnectionContext>& ctx,
Expand All @@ -202,6 +215,7 @@ class ConnectionPoller {
std::mutex wakeup_mutex_;
std::condition_variable wakeup_cv_;
std::thread pollerThread_;
std::atomic<bool> pollerThreadStarted_{false};

std::mutex contexts_mutex_;
std::atomic<uint64_t> contexts_version_{0};
Expand Down
9 changes: 9 additions & 0 deletions mooncake-pg/include/mooncake_backend.h
Original file line number Diff line number Diff line change
Expand Up @@ -150,7 +150,14 @@ class MooncakeBackend final : public ::c10d::Backend {

void recoverRanks(const std::vector<int>& ranks);

void joinGroup();

private:
void waitForExtensionState();
void publishLocalPeerMetadata();
void setLocalOnlyActiveRanks();
void syncActiveRanksTensor();

static TransferEngine* engine_;
std::shared_ptr<MooncakeWorker> worker_;
static bool engineInitialized_;
Expand All @@ -166,6 +173,7 @@ class MooncakeBackend final : public ::c10d::Backend {
std::shared_ptr<TransferGroupMeta> 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
Expand All @@ -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<ConnectionContext> connection_ctx_;
bool connectionPollerRegistered_{false};
};

} // namespace mooncake
Expand Down
45 changes: 44 additions & 1 deletion mooncake-pg/src/connection_poller.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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)),
Expand Down Expand Up @@ -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);
Expand All @@ -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<std::mutex> 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.
{
Expand Down Expand Up @@ -352,13 +383,17 @@ 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(
getServerNameStoreKey(backendIndex_, pollingRank));
store_->deleteKey(getBufferStoreKey(backendIndex_, pollingRank));
store_->deleteKey(
getExtensionTaskCountStoreKey(backendIndex_, pollingRank));
store_->deleteKey(
getExtensionActiveRanksStoreKey(backendIndex_, pollingRank));

// Reset warmup region
*reinterpret_cast<volatile int32_t*>(
Expand Down Expand Up @@ -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<ConnectionContext>& ctx) {
ensureThreadStarted();
{
std::lock_guard<std::mutex> lock(contexts_mutex_);
contexts_.push_back(ctx);
Expand Down
Loading
Loading