Skip to content
Draft
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
32 changes: 27 additions & 5 deletions tests/core/framework/kv_cache_transfer/pd_topology_guard_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -82,14 +82,14 @@ TEST(PdTopologyGuardTest, TryGetPdTopoReturnTopo) {
EXPECT_TRUE(reason.empty());
}

TEST(PdTopologyGuardTest, HeteroTopoNeedMla) {
const InstanceInfo local_info = make_info(2, {0, 1, 2, 3});
const InstanceInfo remote_info = make_info(1, {0, 1, 2, 3});
TEST(PdTopologyGuardTest, HeteroPrefillTpTwoDecodeTpOneAllowsNonMlaPush) {
const InstanceInfo local_info = make_info(1, {0, 1});
const InstanceInfo remote_info = make_info(1, {2});

const PdTopoResult result =
check_pd_topo(local_info, remote_info, "PUSH", false);
EXPECT_EQ(result.status, PdTopoStatus::DENY_HETERO);
EXPECT_EQ(result.reason, "hetero pd requires enable_mla=true");
EXPECT_EQ(result.status, PdTopoStatus::ALLOW_HETERO);
EXPECT_TRUE(result.reason.empty());
}

TEST(PdTopologyGuardTest, HeteroTopoNeedPushKv) {
Expand All @@ -112,6 +112,28 @@ TEST(PdTopologyGuardTest, HeteroTopoAllowOnPushMla) {
EXPECT_TRUE(result.reason.empty());
}

TEST(PdTopologyGuardTest, NonMlaHeteroTopoRequiresEqualDpSize) {
const InstanceInfo local_info = make_info(2, {0, 1, 2, 3});
const InstanceInfo remote_info = make_info(1, {4});

const PdTopoResult result =
check_pd_topo(local_info, remote_info, "PUSH", false);
EXPECT_EQ(result.status, PdTopoStatus::DENY_HETERO);
EXPECT_EQ(result.reason, "non-mla hetero pd requires equal dp_size");
}

TEST(PdTopologyGuardTest, NonMlaHeteroTopoRequiresPrefillTpMultiple) {
const InstanceInfo local_info = make_info(1, {0, 1, 2});
const InstanceInfo remote_info = make_info(1, {3, 4});

const PdTopoResult result =
check_pd_topo(local_info, remote_info, "PUSH", false);
EXPECT_EQ(result.status, PdTopoStatus::DENY_HETERO);
EXPECT_EQ(result.reason,
"non-mla hetero pd requires prefill tp_size divisible by decode "
"tp_size");
}

TEST(PdTopologyGuardTest, CheckPdTopoRejectInvalidLocalTopo) {
const InstanceInfo local_info = make_info(0, {0, 1, 2, 3});
const InstanceInfo remote_info = make_info(1, {0, 1, 2, 3});
Expand Down
18 changes: 18 additions & 0 deletions tests/core/framework/kv_cache_transfer/push_route_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -71,4 +71,22 @@ TEST(PushRouteTest, DstDpRankOffsetApplied) {
EXPECT_EQ(dst_ranks, expect_ranks);
}

TEST(PushRouteTest, DecodeTpOneLinksEveryPrefillTpRank) {
const std::vector<int32_t> src_tp_ranks = get_src_tp_ranks(0, 2, 1);
const std::vector<int32_t> expect_src_tp_ranks = {0, 1};
EXPECT_EQ(src_tp_ranks, expect_src_tp_ranks);
}

TEST(PushRouteTest, DecodeRankLinksMatchingPrefillOwners) {
const std::vector<int32_t> src_tp_ranks = get_src_tp_ranks(1, 8, 2);
const std::vector<int32_t> expect_src_tp_ranks = {1, 3, 5, 7};
EXPECT_EQ(src_tp_ranks, expect_src_tp_ranks);
}

TEST(PushRouteTest, LargerDecodeTpKeepsRoundRobinSource) {
const std::vector<int32_t> src_tp_ranks = get_src_tp_ranks(5, 2, 8);
const std::vector<int32_t> expect_src_tp_ranks = {1};
EXPECT_EQ(src_tp_ranks, expect_src_tp_ranks);
}

} // namespace xllm
24 changes: 24 additions & 0 deletions xllm/core/distributed_runtime/comm_channel.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -266,6 +266,30 @@ bool CommChannel::pull_kv_blocks(
return !cntl.Failed() && s.ok();
}

bool CommChannel::pull_hetero_kv_blocks(
const std::vector<uint64_t>& src_cluster_ids,
const std::vector<std::string>& src_addrs,
const std::vector<uint64_t>& src_blocks,
const std::vector<uint64_t>& dst_blocks,
const std::vector<uint64_t>& src_linear_state_ids,
const std::vector<uint64_t>& dst_linear_state_ids) {
proto::PullKVCacheRequest request;
request.set_hetero_merge(true);
ADD_VECTOR_TO_PROTO(request.mutable_src_cluster_ids(), src_cluster_ids);
ADD_VECTOR_TO_PROTO(request.mutable_src_addrs(), src_addrs);
ADD_VECTOR_TO_PROTO(request.mutable_src_blocks(), src_blocks);
ADD_VECTOR_TO_PROTO(request.mutable_dst_blocks(), dst_blocks);
ADD_VECTOR_TO_PROTO(request.mutable_src_linear_state_ids(),
src_linear_state_ids);
ADD_VECTOR_TO_PROTO(request.mutable_dst_linear_state_ids(),
dst_linear_state_ids);

proto::Status s;
brpc::Controller cntl;
stub_->PullKVCache(&cntl, &request, &s, nullptr);
return !cntl.Failed() && s.ok();
}

void CommChannel::execute_model_async(
const ForwardInput& input,
folly::Promise<std::optional<RawForwardOutput>>& promise) {
Expand Down
8 changes: 8 additions & 0 deletions xllm/core/distributed_runtime/comm_channel.h
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,14 @@ class CommChannel {
const std::vector<uint64_t>& src_linear_state_ids = {},
const std::vector<uint64_t>& dst_linear_state_ids = {});

virtual bool pull_hetero_kv_blocks(
const std::vector<uint64_t>& src_cluster_ids,
const std::vector<std::string>& src_addrs,
const std::vector<uint64_t>& src_blocks,
const std::vector<uint64_t>& dst_blocks,
const std::vector<uint64_t>& src_linear_state_ids = {},
const std::vector<uint64_t>& dst_linear_state_ids = {});

virtual void execute_model_async(
const ForwardInput& input,
folly::Promise<std::optional<RawForwardOutput>>& promise);
Expand Down
14 changes: 14 additions & 0 deletions xllm/core/distributed_runtime/engine.h
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,20 @@ class Engine {
return false;
};

virtual bool pull_hetero_kv_blocks(
const int32_t src_dp_size,
const int32_t src_dp_rank,
const std::vector<uint64_t>& src_cluster_ids,
const std::vector<std::string>& src_addrs,
const std::vector<uint64_t>& src_blocks,
const int32_t dst_dp_rank,
const std::vector<uint64_t>& dst_blocks,
const std::vector<uint64_t>& src_linear_state_ids = {},
const std::vector<uint64_t>& dst_linear_state_ids = {}) {
NOT_IMPLEMENTED();
return false;
};

virtual std::vector<folly::SemiFuture<uint32_t>> transfer_kv_blocks(
const uint32_t dp_rank,
const std::vector<BlockTransferInfo>& block_transfer_info) {
Expand Down
121 changes: 91 additions & 30 deletions xllm/core/distributed_runtime/llm_engine.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@ limitations under the License.
#include "framework/kv_cache/kv_cache_estimation.h"
#include "framework/kv_cache/kv_cache_shape.h"
#include "framework/kv_cache/kv_cache_utils.h"
#include "framework/kv_cache_transfer/push_route.h"
#include "framework/model/model_args.h"
#include "framework/model_loader.h"
#include "framework/xtensor/page_allocator.h"
Expand Down Expand Up @@ -701,6 +702,55 @@ bool LLMEngine::pull_kv_blocks(
return true;
}

bool LLMEngine::pull_hetero_kv_blocks(
const int32_t src_dp_size,
const int32_t src_dp_rank,
const std::vector<uint64_t>& src_cluster_ids,
const std::vector<std::string>& src_addrs,
const std::vector<uint64_t>& src_blocks,
const int32_t dst_dp_rank,
const std::vector<uint64_t>& dst_blocks,
const std::vector<uint64_t>& src_linear_state_ids,
const std::vector<uint64_t>& dst_linear_state_ids) {
if (src_dp_size <= 0 || src_dp_rank < 0 || src_dp_rank >= src_dp_size ||
src_cluster_ids.size() != src_addrs.size() ||
src_cluster_ids.size() % static_cast<size_t>(src_dp_size) != 0) {
LOG(ERROR) << "Invalid heterogeneous KV pull topology.";
return false;
}
const int32_t src_tp_size =
static_cast<int32_t>(src_cluster_ids.size()) / src_dp_size;
const int32_t dst_tp_size = static_cast<int32_t>(dp_local_tp_size_);
if (src_tp_size < dst_tp_size || src_tp_size % dst_tp_size != 0) {
LOG(ERROR) << "Unsupported heterogeneous KV pull ratio: prefill_tp_size="
<< src_tp_size << ", decode_tp_size=" << dst_tp_size;
return false;
}

std::vector<bool> results;
results.reserve(dst_tp_size);
for (int32_t dst_tp_rank = 0; dst_tp_rank < dst_tp_size; ++dst_tp_rank) {
std::vector<uint64_t> worker_src_cluster_ids;
std::vector<std::string> worker_src_addrs;
for (int32_t src_tp_rank :
get_src_tp_ranks(dst_tp_rank, src_tp_size, dst_tp_size)) {
const int32_t src_worker_rank = src_dp_rank * src_tp_size + src_tp_rank;
worker_src_cluster_ids.push_back(src_cluster_ids[src_worker_rank]);
worker_src_addrs.push_back(src_addrs[src_worker_rank]);
}
const int32_t dst_worker_rank = dst_dp_rank * dst_tp_size + dst_tp_rank;
results.push_back(worker_clients_[dst_worker_rank]->pull_hetero_kv_blocks(
worker_src_cluster_ids,
worker_src_addrs,
src_blocks,
dst_blocks,
src_linear_state_ids,
dst_linear_state_ids));
}
return std::all_of(
results.begin(), results.end(), [](bool result) { return result; });
}

std::vector<folly::SemiFuture<uint32_t>> LLMEngine::transfer_kv_blocks(
const uint32_t dp_rank,
const std::vector<BlockTransferInfo>& block_transfer_info) {
Expand Down Expand Up @@ -786,38 +836,43 @@ bool LLMEngine::link_cluster(const std::vector<uint64_t>& cluster_ids,
const std::vector<uint16_t>& ports,
const int32_t src_dp_size,
const int32_t src_kv_split_size) {
// Each D worker connects to all P workers that share the same TP rank.
const int32_t src_world_size = static_cast<int32_t>(cluster_ids.size());

// Each D worker connects to every P worker that routes KV blocks to its TP
// rank. When P TP is larger, multiple P ranks share one D-side owner.
// P layout: rank = dp_i * src_cp_tp_size + split_j * src_tp_size + tp_rank
// D workers cycle through tp_rank in [0, src_tp_size) round-robin.
// Requires: D-side dp_local_tp_size_ == src_tp_size.
int32_t src_world_size = static_cast<int32_t>(cluster_ids.size());
int32_t src_cp_tp_size = src_world_size / src_dp_size;
int32_t src_tp_size = src_cp_tp_size / src_kv_split_size;
int32_t src_dp_worker_index = 0;

std::vector<folly::SemiFuture<bool>> futures;
futures.reserve(worker_clients_num_);
for (size_t worker_rank = 0; worker_rank < worker_clients_num_;
++worker_rank) {
std::vector<uint64_t> target_cluster_ids;
std::vector<std::string> target_addrs;
std::vector<uint16_t> target_ports;
target_cluster_ids.reserve(src_dp_size * src_kv_split_size);
target_addrs.reserve(src_dp_size * src_kv_split_size);
target_ports.reserve(src_dp_size * src_kv_split_size);
const int32_t dst_tp_rank =
static_cast<int32_t>(worker_rank % dp_local_tp_size_);
const std::vector<int32_t> src_tp_ranks =
get_src_tp_ranks(dst_tp_rank, src_tp_size, dp_local_tp_size_);
const size_t endpoint_count = static_cast<size_t>(src_dp_size) *
static_cast<size_t>(src_kv_split_size) *
src_tp_ranks.size();
target_cluster_ids.reserve(endpoint_count);
target_addrs.reserve(endpoint_count);
target_ports.reserve(endpoint_count);

for (int32_t dp_i = 0; dp_i < src_dp_size; ++dp_i) {
for (int32_t split_j = 0; split_j < src_kv_split_size; ++split_j) {
int32_t p_idx =
dp_i * src_cp_tp_size + split_j * src_tp_size + src_dp_worker_index;
target_cluster_ids.emplace_back(cluster_ids[p_idx]);
target_addrs.emplace_back(addrs[p_idx]);
target_ports.emplace_back(ports[p_idx]);
for (int32_t src_tp_rank : src_tp_ranks) {
const int32_t p_idx =
dp_i * src_cp_tp_size + split_j * src_tp_size + src_tp_rank;
target_cluster_ids.emplace_back(cluster_ids[p_idx]);
target_addrs.emplace_back(addrs[p_idx]);
target_ports.emplace_back(ports[p_idx]);
}
}
}

src_dp_worker_index = (src_dp_worker_index + 1) % src_tp_size;

folly::Promise<bool> promise;
auto future = promise.getSemiFuture();
link_threadpool_->schedule(
Expand Down Expand Up @@ -848,35 +903,41 @@ bool LLMEngine::unlink_cluster(const std::vector<uint64_t>& cluster_ids,
const std::vector<uint16_t>& ports,
const int32_t src_dp_size,
const int32_t src_kv_split_size) {
// Symmetric to link_cluster; uses the same rank mapping.
int32_t src_world_size = static_cast<int32_t>(cluster_ids.size());
const int32_t src_world_size = static_cast<int32_t>(cluster_ids.size());

// Symmetric to link_cluster; uses the same owner mapping.
int32_t src_cp_tp_size = src_world_size / src_dp_size;
int32_t src_tp_size = src_cp_tp_size / src_kv_split_size;
int32_t src_dp_worker_index = 0;

std::vector<folly::SemiFuture<bool>> futures;
futures.reserve(worker_clients_num_);
for (size_t worker_rank = 0; worker_rank < worker_clients_num_;
++worker_rank) {
std::vector<uint64_t> target_cluster_ids;
std::vector<std::string> target_addrs;
std::vector<uint16_t> target_ports;
target_cluster_ids.reserve(src_dp_size * src_kv_split_size);
target_addrs.reserve(src_dp_size * src_kv_split_size);
target_ports.reserve(src_dp_size * src_kv_split_size);
const int32_t dst_tp_rank =
static_cast<int32_t>(worker_rank % dp_local_tp_size_);
const std::vector<int32_t> src_tp_ranks =
get_src_tp_ranks(dst_tp_rank, src_tp_size, dp_local_tp_size_);
const size_t endpoint_count = static_cast<size_t>(src_dp_size) *
static_cast<size_t>(src_kv_split_size) *
src_tp_ranks.size();
target_cluster_ids.reserve(endpoint_count);
target_addrs.reserve(endpoint_count);
target_ports.reserve(endpoint_count);

for (int32_t dp_i = 0; dp_i < src_dp_size; ++dp_i) {
for (int32_t split_j = 0; split_j < src_kv_split_size; ++split_j) {
int32_t p_idx =
dp_i * src_cp_tp_size + split_j * src_tp_size + src_dp_worker_index;
target_cluster_ids.emplace_back(cluster_ids[p_idx]);
target_addrs.emplace_back(addrs[p_idx]);
target_ports.emplace_back(ports[p_idx]);
for (int32_t src_tp_rank : src_tp_ranks) {
const int32_t p_idx =
dp_i * src_cp_tp_size + split_j * src_tp_size + src_tp_rank;
target_cluster_ids.emplace_back(cluster_ids[p_idx]);
target_addrs.emplace_back(addrs[p_idx]);
target_ports.emplace_back(ports[p_idx]);
}
}
}

src_dp_worker_index = (src_dp_worker_index + 1) % src_tp_size;

folly::Promise<bool> promise;
auto future = promise.getSemiFuture();
link_threadpool_->schedule(
Expand Down
11 changes: 11 additions & 0 deletions xllm/core/distributed_runtime/llm_engine.h
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,17 @@ class LLMEngine : public Engine {
const std::vector<uint64_t>& src_linear_state_ids = {},
const std::vector<uint64_t>& dst_linear_state_ids = {}) override;

bool pull_hetero_kv_blocks(
const int32_t src_dp_size,
const int32_t src_dp_rank,
const std::vector<uint64_t>& src_cluster_ids,
const std::vector<std::string>& src_addrs,
const std::vector<uint64_t>& src_blocks,
const int32_t dst_dp_rank,
const std::vector<uint64_t>& dst_blocks,
const std::vector<uint64_t>& src_linear_state_ids = {},
const std::vector<uint64_t>& dst_linear_state_ids = {}) override;

std::vector<folly::SemiFuture<uint32_t>> transfer_kv_blocks(
const uint32_t dp_rank,
const std::vector<BlockTransferInfo>& block_transfer_info) override;
Expand Down
15 changes: 15 additions & 0 deletions xllm/core/distributed_runtime/remote_worker.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -275,6 +275,21 @@ folly::SemiFuture<bool> RemoteWorker::pull_kv_blocks_async(
return future;
}

bool RemoteWorker::pull_hetero_kv_blocks(
const std::vector<uint64_t>& src_cluster_ids,
const std::vector<std::string>& src_addrs,
const std::vector<uint64_t>& src_blocks,
const std::vector<uint64_t>& dst_blocks,
const std::vector<uint64_t>& src_linear_state_ids,
const std::vector<uint64_t>& dst_linear_state_ids) {
return channel_->pull_hetero_kv_blocks(src_cluster_ids,
src_addrs,
src_blocks,
dst_blocks,
src_linear_state_ids,
dst_linear_state_ids);
}

folly::SemiFuture<uint32_t> RemoteWorker::transfer_kv_blocks(
const std::vector<BlockTransferInfo>& block_transfer_info) {
folly::Promise<uint32_t> promise;
Expand Down
8 changes: 8 additions & 0 deletions xllm/core/distributed_runtime/remote_worker.h
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,14 @@ class RemoteWorker : public WorkerClient {
const std::vector<uint64_t>& src_linear_state_ids = {},
const std::vector<uint64_t>& dst_linear_state_ids = {}) override;

virtual bool pull_hetero_kv_blocks(
const std::vector<uint64_t>& src_cluster_ids,
const std::vector<std::string>& src_addrs,
const std::vector<uint64_t>& src_blocks,
const std::vector<uint64_t>& dst_blocks,
const std::vector<uint64_t>& src_linear_state_ids = {},
const std::vector<uint64_t>& dst_linear_state_ids = {}) override;

// prepare input request
virtual ForwardInput prepare_inputs(Batch& batch) override;

Expand Down
Loading
Loading