fix(mooncake): honour MOONCAKE_PROTOCOL so EFA hardware can select efa transport (#25083)

Co-authored-by: whn09 <whn09@users.noreply.github.com>
This commit is contained in:
王鹤男
2026-05-29 14:21:15 +08:00
committed by GitHub
co-authored by whn09
parent 73c99e3361
commit 5850aa14c3
3 changed files with 18 additions and 18 deletions
@@ -183,21 +183,18 @@ class MooncakeTransferEngine:
"""Initialize the mooncake instance.""" """Initialize the mooncake instance."""
if envs.ENABLE_ASCEND_TRANSFER_WITH_MOONCAKE.get(): if envs.ENABLE_ASCEND_TRANSFER_WITH_MOONCAKE.get():
npu_phy_id = envs.ASCEND_NPU_PHY_ID.get() npu_phy_id = envs.ASCEND_NPU_PHY_ID.get()
if npu_phy_id == -1: suffix = self.gpu_id if npu_phy_id == -1 else npu_phy_id
hostname += f":{get_free_port()}:npu_{self.gpu_id}" hostname += f":{get_free_port()}:npu_{suffix}"
protocol = "ascend"
else: else:
hostname += f":{get_free_port()}:npu_{npu_phy_id}" # MOONCAKE_PROTOCOL selects the transport (rdma | efa | tcp | ...).
# Default is "rdma"; set MOONCAKE_PROTOCOL=efa on AWS EFA hardware.
protocol = envs.MOONCAKE_PROTOCOL.get()
ret_value = self.engine.initialize( ret_value = self.engine.initialize(
hostname, hostname,
"P2PHANDSHAKE", "P2PHANDSHAKE",
"ascend", protocol,
device_name if device_name is not None else "",
)
else:
ret_value = self.engine.initialize(
hostname,
"P2PHANDSHAKE",
"rdma",
device_name if device_name is not None else "", device_name if device_name is not None else "",
) )
if ret_value != 0: if ret_value != 0:
+1 -1
View File
@@ -356,7 +356,7 @@ class Envs:
MOONCAKE_LOCAL_HOSTNAME = EnvStr("localhost") MOONCAKE_LOCAL_HOSTNAME = EnvStr("localhost")
MOONCAKE_TE_META_DATA_SERVER = EnvStr("P2PHANDSHAKE") MOONCAKE_TE_META_DATA_SERVER = EnvStr("P2PHANDSHAKE")
MOONCAKE_GLOBAL_SEGMENT_SIZE = EnvStr("4gb") MOONCAKE_GLOBAL_SEGMENT_SIZE = EnvStr("4gb")
MOONCAKE_PROTOCOL = EnvStr("tcp") MOONCAKE_PROTOCOL = EnvStr("rdma")
MOONCAKE_DEVICE = EnvStr("") MOONCAKE_DEVICE = EnvStr("")
MOONCAKE_MASTER_METRICS_PORT = EnvInt(9003) MOONCAKE_MASTER_METRICS_PORT = EnvInt(9003)
MOONCAKE_CHECK_SERVER = EnvBool(False) MOONCAKE_CHECK_SERVER = EnvBool(False)
@@ -925,7 +925,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.remote_instance_transfer_engine = TransferEngine() self.remote_instance_transfer_engine = TransferEngine()
local_ip = get_local_ip_auto() local_ip = get_local_ip_auto()
self.remote_instance_transfer_engine.initialize( self.remote_instance_transfer_engine.initialize(
local_ip, "P2PHANDSHAKE", "rdma", envs.MOONCAKE_DEVICE.get() local_ip,
"P2PHANDSHAKE",
envs.MOONCAKE_PROTOCOL.get(),
envs.MOONCAKE_DEVICE.get(),
) )
self.remote_instance_transfer_engine_session_id = NetworkAddress( self.remote_instance_transfer_engine_session_id = NetworkAddress(
local_ip, self.remote_instance_transfer_engine.get_rpc_port() local_ip, self.remote_instance_transfer_engine.get_rpc_port()