[PD] feat: refactor custom mem pool and add barex pd support (#12332)

This commit is contained in:
Teng Ma
2025-11-09 18:51:34 +08:00
committed by GitHub
parent 877cb52840
commit b8ac4fcb51
5 changed files with 162 additions and 49 deletions
@@ -28,14 +28,12 @@ from sglang.srt.disaggregation.common.utils import (
group_concurrent_contiguous,
)
from sglang.srt.disaggregation.mooncake.transfer_engine import MooncakeTransferEngine
from sglang.srt.disaggregation.mooncake.utils import (
check_mooncake_custom_mem_pool_enabled,
)
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import (
format_tcp_address,
get_bool_env_var,
get_int_env_var,
is_valid_ipv6_address,
)
from sglang.srt.utils import format_tcp_address, get_int_env_var, is_valid_ipv6_address
logger = logging.getLogger(__name__)
@@ -201,8 +199,8 @@ class MooncakeKVManager(CommonKVManager):
"SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT", 300
)
self.enable_custom_mem_pool = get_bool_env_var(
"SGLANG_MOONCAKE_CUSTOM_MEM_POOL", "false"
self.enable_custom_mem_pool, self.custom_mem_pool_type = (
check_mooncake_custom_mem_pool_enabled()
)
elif self.disaggregation_mode == DisaggregationMode.DECODE:
self.heartbeat_failures = {}
@@ -549,7 +547,7 @@ class MooncakeKVManager(CommonKVManager):
dst_aux_ptrs: list[int],
):
# TODO(shangming): Fix me when nvlink_transport of Mooncake is bug-free
if self.enable_custom_mem_pool:
if self.enable_custom_mem_pool and self.custom_mem_pool_type == "NVLINK":
return self.send_aux_tcp(req, prefill_aux_index, dst_aux_ptrs)
transfer_blocks = []