[PD][MoRI] Drive KV transfers with a sharded synchronous worker pool (#26922)

This commit is contained in:
Niko Ma
2026-06-08 00:49:25 -07:00
committed by GitHub
parent 0d0254c9de
commit 18d728967a
4 changed files with 235 additions and 112 deletions
+1 -1
View File
@@ -104,7 +104,7 @@ ARG ENABLE_MORI=0
ARG NIC_BACKEND=none ARG NIC_BACKEND=none
ARG MORI_REPO="https://github.com/ROCm/mori.git" ARG MORI_REPO="https://github.com/ROCm/mori.git"
ARG MORI_COMMIT="96ffa169710f214e76e07abe5008d686fe54522b" ARG MORI_COMMIT="d87651c998296c5bddbeefc4cc525b58663b1636"
# AMD AINIC apt repo settings # AMD AINIC apt repo settings
ARG AINIC_VERSION=1.117.5-a-38 ARG AINIC_VERSION=1.117.5-a-38
+208 -96
View File
@@ -23,6 +23,7 @@ from mori.io import (
MemoryLocationType, MemoryLocationType,
PollCqMode, PollCqMode,
RdmaBackendConfig, RdmaBackendConfig,
StatusCode,
) )
from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll from sglang.srt.disaggregation.base.conn import KVArgs, KVPoll
@@ -34,13 +35,14 @@ from sglang.srt.disaggregation.common.conn import (
) )
from sglang.srt.disaggregation.common.utils import ( from sglang.srt.disaggregation.common.utils import (
AuxDataCodec, AuxDataCodec,
FastQueue,
group_concurrent_contiguous, group_concurrent_contiguous,
pack_int_lists, pack_int_lists,
unpack_int_lists, unpack_int_lists,
) )
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.environ import envs
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import get_int_env_var
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -267,6 +269,16 @@ class TransferTarget:
peer_info: KVArgsRegisterInfo peer_info: KVArgsRegisterInfo
@dataclasses.dataclass
class _TransferChunk:
sender: "MoriKVSender"
kv_indices: npt.NDArray[np.int32]
index_slice: slice
is_last_chunk: bool
aux_index: Optional[int]
normalized_state: Optional[List[Optional[npt.NDArray[np.int32]]]]
class MoriKVManager(CommonKVManager): class MoriKVManager(CommonKVManager):
AUX_DATA_HEADER = b"AUX_DATA" AUX_DATA_HEADER = b"AUX_DATA"
@@ -286,13 +298,25 @@ class MoriKVManager(CommonKVManager):
self.transfer_lock = threading.Lock() self.transfer_lock = threading.Lock()
self._zmq_ctx = zmq.Context() self._zmq_ctx = zmq.Context()
self._socket_local = threading.local() self._socket_local = threading.local()
# Send CPU-resident AUX data via RDMA instead of ZMQ TCP. self._send_aux_rdma = envs.SGLANG_MORI_SEND_AUX_RDMA.get()
# Default: TCP. Set SGLANG_MORI_SEND_AUX_RDMA=1 to use RDMA.
self._send_aux_rdma = os.environ.get(
"SGLANG_MORI_SEND_AUX_RDMA", ""
).lower() in ("1", "true")
self._register_local_buffers() self._register_local_buffers()
if self.disaggregation_mode == DisaggregationMode.PREFILL: if self.disaggregation_mode == DisaggregationMode.PREFILL:
self._num_shards = max(1, envs.SGLANG_MORI_TRANSFER_SHARDS.get())
self._transfer_queues: List[FastQueue] = [
FastQueue() for _ in range(self._num_shards)
]
self._wait_poll_ms = envs.SGLANG_MORI_WAIT_POLL_MS.get()
self._transfer_timeout_ms = envs.SGLANG_MORI_TRANSFER_TIMEOUT_MS.get()
for shard, queue in enumerate(self._transfer_queues):
threading.Thread(
target=self._transfer_worker,
args=(queue,),
daemon=True,
name=(
f"mori-xfer-dp{self.system_dp_rank}-"
f"tp{self.attn_tp_rank}-s{shard}"
),
).start()
self._start_bootstrap_thread() self._start_bootstrap_thread()
elif self.disaggregation_mode == DisaggregationMode.DECODE: elif self.disaggregation_mode == DisaggregationMode.DECODE:
self.room_to_bootstrap_addr: Dict[int, str] = {} self.room_to_bootstrap_addr: Dict[int, str] = {}
@@ -315,24 +339,9 @@ class MoriKVManager(CommonKVManager):
engine = IOEngine(engine_key, config) engine = IOEngine(engine_key, config)
poll_mode = PollCqMode.POLLING poll_mode = PollCqMode.POLLING
# Number of RDMA Queue Pairs (QPs) used per transfer operation. qp_per_transfer = envs.SGLANG_MORI_QP_PER_TRANSFER.get()
# Higher values can increase parallelism and bandwidth utilization. post_batch_size = envs.SGLANG_MORI_POST_BATCH_SIZE.get()
# Default: 4 num_worker_threads = envs.SGLANG_MORI_NUM_WORKERS.get()
qp_per_transfer = get_int_env_var("SGLANG_MORI_QP_PER_TRANSFER", 4)
# Number of RDMA work requests posted in a single batch to each QP.
# Larger batch sizes reduce per-operation overhead and improve throughput
# at the cost of higher latency. Use -1 for automatic sizing based on
# the number of merged work requests and available endpoints.
# Default: -1 (automatic)
post_batch_size = get_int_env_var("SGLANG_MORI_POST_BATCH_SIZE", -1)
# Number of worker threads in the RDMA executor thread pool.
# Each worker handles RDMA operations on a separate CPU core (with affinity).
# More workers can improve parallelism for large batch transfers across
# multiple QPs, but excessive threads may cause contention.
# Default: 4
num_worker_threads = get_int_env_var("SGLANG_MORI_NUM_WORKERS", 4)
rdma_cfg = RdmaBackendConfig( rdma_cfg = RdmaBackendConfig(
qp_per_transfer, qp_per_transfer,
@@ -400,6 +409,34 @@ class MoriKVManager(CommonKVManager):
return return
super().update_status(bootstrap_room, status) super().update_status(bootstrap_room, status)
def enqueue_transfer(self, task: _TransferChunk) -> None:
self._transfer_queues[task.sender.bootstrap_room % self._num_shards].put(task)
def _transfer_worker(self, queue: FastQueue) -> None:
while True:
task = queue.get()
try:
task.sender._run_chunk(task)
except Exception as exc:
failure_reason = f"transfer worker raised: {exc!r}"
try:
logger.exception(
"Mori transfer worker failed for room %s",
task.sender.bootstrap_room,
)
except Exception:
pass
try:
task.sender._fail_from_worker(failure_reason)
except Exception:
try:
logger.exception(
"Mori transfer worker failover failed for room %s",
task.sender.bootstrap_room,
)
except Exception:
pass
def _connect_threadsafe(self, endpoint: str, is_ipv6: bool = False): def _connect_threadsafe(self, endpoint: str, is_ipv6: bool = False):
"""Thread-local ZMQ socket cache with shared Context. """Thread-local ZMQ socket cache with shared Context.
@@ -799,11 +836,13 @@ class MoriKVManager(CommonKVManager):
kv_item_len = self.kv_args.kv_item_lens[0] kv_item_len = self.kv_args.kv_item_lens[0]
if self.is_mla_backend: if self.is_mla_backend:
layer_plan = self._build_contiguous_transfer_plan(grouped_plan, kv_item_len)
src_descs, dst_descs, layers_current_pp_stage = ( src_descs, dst_descs, layers_current_pp_stage = (
self._get_mla_mem_desc_slices(peer_info.dst_kv_mem_descs) self._get_mla_mem_desc_slices(peer_info.dst_kv_mem_descs)
) )
for layer_id in range(layers_current_pp_stage): for layer_id in range(layers_current_pp_stage):
layer_plan = self._build_contiguous_transfer_plan(
grouped_plan, self.kv_args.kv_item_lens[layer_id]
)
statuses.extend( statuses.extend(
self._submit_batch_transfer_plan( self._submit_batch_transfer_plan(
src_descs[layer_id], src_descs[layer_id],
@@ -1273,12 +1312,6 @@ class MoriKVManager(CommonKVManager):
) )
return result_statuses, target_infos_snapshot return result_statuses, target_infos_snapshot
if is_last_chunk:
with self.transfer_lock:
# Keep transfer_infos alive until sender.clear() so abort/failure
# paths can still recover notification targets after posting.
self.update_status(bootstrap_room, KVPoll.Success)
return result_statuses, target_infos_snapshot return result_statuses, target_infos_snapshot
@@ -1294,10 +1327,12 @@ class MoriKVSender(CommonKVSender):
super().__init__(mgr, bootstrap_addr, bootstrap_room, dest_tp_ranks, pp_rank) super().__init__(mgr, bootstrap_addr, bootstrap_room, dest_tp_ranks, pp_rank)
self.transfer_statuses: List[TransferStatus] = [] self.transfer_statuses: List[TransferStatus] = []
self.pending_infos: Optional[List[TransferInfo]] = None self.pending_infos: Optional[List[TransferInfo]] = None
self.sent_last_chunk = False
self.conclude_state: Optional[KVPoll] = None self.conclude_state: Optional[KVPoll] = None
self.status_notified = False self.status_notified = False
self.init_time = time.time() self.init_time = time.time()
self._notify_lock = threading.Lock()
self._notified_status: Optional[KVPoll] = None
self._notified_reason: Optional[str] = None
def send( def send(
self, self,
@@ -1315,20 +1350,17 @@ class MoriKVSender(CommonKVSender):
if is_last_chunk if is_last_chunk
else None else None
) )
statuses, infos = self.kv_mgr.add_transfer_request( self._record_transfer_indices(kv_indices, state_indices)
self.bootstrap_room, self.kv_mgr.enqueue_transfer(
kv_indices, _TransferChunk(
index_slice, sender=self,
is_last_chunk, kv_indices=kv_indices,
aux_index=self.aux_index if is_last_chunk else None, index_slice=index_slice,
state_indices=normalized_state, is_last_chunk=is_last_chunk,
aux_index=self.aux_index if is_last_chunk else None,
normalized_state=normalized_state,
)
) )
self.transfer_statuses.extend(statuses)
self._record_transfer_indices(kv_indices, None)
if infos is not None:
self.pending_infos = infos
if is_last_chunk:
self.sent_last_chunk = True
self._maybe_finalize_if_room_failed() self._maybe_finalize_if_room_failed()
def _maybe_finalize_if_room_failed(self) -> None: def _maybe_finalize_if_room_failed(self) -> None:
@@ -1337,53 +1369,103 @@ class MoriKVSender(CommonKVSender):
if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed: if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed:
self._finalize_failure() self._finalize_failure()
def _run_chunk(self, task: _TransferChunk) -> None:
if self.conclude_state is not None:
return
if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed:
self._finalize_failure()
return
statuses, infos = self.kv_mgr.add_transfer_request(
self.bootstrap_room,
task.kv_indices,
task.index_slice,
task.is_last_chunk,
aux_index=task.aux_index,
state_indices=task.normalized_state,
)
self.transfer_statuses.extend(statuses)
if infos is not None:
self.pending_infos = infos
if self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed:
self._finalize_failure()
return
rc = self._wait_chunk(statuses)
if self.conclude_state is not None:
return
if rc != StatusCode.SUCCESS:
self._finalize_failure(self._collect_failure_reason())
return
if task.is_last_chunk:
self._notify_decode(KVPoll.Success)
with self._notify_lock:
if self.conclude_state is None:
self.conclude_state = self._notified_status
if self._notified_status == KVPoll.Success:
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Success)
def _wait_chunk(self, statuses: List[TransferStatus]) -> StatusCode:
if not statuses:
return StatusCode.SUCCESS
start = time.perf_counter()
sla_ms = self.kv_mgr._transfer_timeout_ms
sla_tripped = False
while True:
rc = self.kv_mgr.engine.wait_all(
statuses, timeout_ms=self.kv_mgr._wait_poll_ms
)
if rc != StatusCode.IN_PROGRESS:
return rc
if (
sla_ms > 0
and not sla_tripped
and (time.perf_counter() - start) * 1000 >= sla_ms
):
sla_tripped = True
self._finalize_failure(f"KV transfer exceeded SLA {sla_ms}ms")
def _fail_from_worker(self, reason: str) -> None:
self._finalize_failure(reason)
def poll(self) -> KVPoll: def poll(self) -> KVPoll:
if self.conclude_state is not None: if self.conclude_state is not None:
return self.conclude_state return self.conclude_state
if self.bootstrap_room not in self.kv_mgr.request_status: if self.bootstrap_room not in self.kv_mgr.request_status:
self._finalize_failure() sent_status, _ = self._finalize_failure()
return KVPoll.Failed return sent_status
status = self.kv_mgr.check_status(self.bootstrap_room) status = self.kv_mgr.check_status(self.bootstrap_room)
if status == KVPoll.Bootstrapping: if status == KVPoll.Bootstrapping:
timeout_result = self._check_bootstrap_timeout() elapsed = time.time() - self.init_time
if timeout_result is not None: if elapsed >= self.kv_mgr.bootstrap_timeout:
self._finalize_failure() logger.warning_once(
return KVPoll.Failed "Some requests timed out when bootstrapping, "
"which means prefill instances fail to receive the KV indices from the decode instance of this request. "
"If a greater mean TTFT is acceptable, you can 'export SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT=600' (10 minutes) to relax the timeout condition. "
)
reason = (
f"Request {self.bootstrap_room} timed out after {elapsed:.1f}s "
"in KVPoll.Bootstrapping"
)
sent_status, _ = self._finalize_failure(reason)
return sent_status
return status return status
if status == KVPoll.Failed: if status == KVPoll.Failed:
self._finalize_failure() sent_status, _ = self._finalize_failure()
return KVPoll.Failed return sent_status
if status == KVPoll.Success and self.kv_mgr.is_dummy_cp_rank: if status == KVPoll.Success:
self.conclude_state = KVPoll.Success self.conclude_state = KVPoll.Success
return KVPoll.Success return KVPoll.Success
transfers_done = self._all_transfers_finished() return status
if transfers_done:
if self._has_transfer_error():
reason = self._collect_failure_reason()
self.kv_mgr.record_failure(self.bootstrap_room, reason)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
self._finalize_failure(reason)
return KVPoll.Failed
self._notify_decode(KVPoll.Success)
self.conclude_state = KVPoll.Success
return KVPoll.Success
return KVPoll.Transferring if status == KVPoll.Success else status
def _all_transfers_finished(self) -> bool:
if not self.sent_last_chunk:
return False
if not self.transfer_statuses:
return True
return all(not status.InProgress() for status in self.transfer_statuses)
def _has_transfer_error(self) -> bool:
return any(status.Failed() for status in self.transfer_statuses)
def _collect_failure_reason(self) -> str: def _collect_failure_reason(self) -> str:
for status in self.transfer_statuses: for status in self.transfer_statuses:
@@ -1391,33 +1473,66 @@ class MoriKVSender(CommonKVSender):
return f"KV transfer failed: {status.Message()}" return f"KV transfer failed: {status.Message()}"
return "KV transfer failed due to unknown reason" return "KV transfer failed due to unknown reason"
def _notify_decode( def _terminalize_locked(
self, status: KVPoll, failure_reason: Optional[str] = None self,
) -> None: status: KVPoll,
reason: Optional[str] = None,
) -> Tuple[KVPoll, Optional[str], Optional[List[TransferInfo]]]:
if self.status_notified: if self.status_notified:
return return self._notified_status, self._notified_reason, None
if status == KVPoll.Success:
with self.kv_mgr.failure_lock:
recorded = self.kv_mgr.failure_records.get(self.bootstrap_room)
if recorded is not None:
status = KVPoll.Failed
reason = recorded
elif self.kv_mgr.request_status.get(self.bootstrap_room) == KVPoll.Failed:
status = KVPoll.Failed
reason = reason or "request marked Failed before notify"
if status == KVPoll.Failed:
with self.kv_mgr.failure_lock:
self.kv_mgr.failure_records.setdefault(
self.bootstrap_room, reason or "KV transfer failed"
)
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
infos = self.pending_infos infos = self.pending_infos
if infos is None: if infos is None:
with self.kv_mgr.transfer_lock: with self.kv_mgr.transfer_lock:
room_infos = self.kv_mgr.transfer_infos.get(self.bootstrap_room) room_infos = self.kv_mgr.transfer_infos.get(self.bootstrap_room)
if room_infos is not None: infos = list(room_infos.values()) if room_infos is not None else None
infos = list(room_infos.values())
self._notified_status = status
self._notified_reason = reason
self.status_notified = True
return status, reason, infos
def _notify_decode(
self, status: KVPoll, failure_reason: Optional[str] = None
) -> Tuple[KVPoll, Optional[str]]:
with self._notify_lock:
emitted_status, emitted_reason, infos = self._terminalize_locked(
status, failure_reason
)
if infos: if infos:
self.kv_mgr.notify_decode_status( self.kv_mgr.notify_decode_status(
infos, self.bootstrap_room, status, failure_reason infos, self.bootstrap_room, emitted_status, emitted_reason
) )
self.status_notified = True return emitted_status, emitted_reason
def _finalize_failure(self, failure_reason: Optional[str] = None) -> None: def _finalize_failure(
if self.conclude_state == KVPoll.Failed: self, failure_reason: Optional[str] = None
return ) -> Tuple[KVPoll, Optional[str]]:
if failure_reason is None: if failure_reason is None:
failure_reason = self.kv_mgr.failure_records.get( with self.kv_mgr.failure_lock:
self.bootstrap_room, "KV transfer failed" failure_reason = self.kv_mgr.failure_records.get(
) self.bootstrap_room, "KV transfer failed"
self._notify_decode(KVPoll.Failed, failure_reason) )
self.conclude_state = KVPoll.Failed sent_status, sent_reason = self._notify_decode(KVPoll.Failed, failure_reason)
self.conclude_state = sent_status
return sent_status, sent_reason
def failure_exception(self): def failure_exception(self):
if self.conclude_state is None: if self.conclude_state is None:
@@ -1430,10 +1545,7 @@ class MoriKVSender(CommonKVSender):
raise RuntimeError(failure_reason) raise RuntimeError(failure_reason)
def abort(self): def abort(self):
self.kv_mgr.record_failure(self.bootstrap_room, "Aborted by AbortReq.") self._finalize_failure("Aborted by AbortReq.")
self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed)
self._notify_decode(KVPoll.Failed, "Aborted by AbortReq.")
self.conclude_state = KVPoll.Failed
class MoriKVReceiver(CommonKVReceiver): class MoriKVReceiver(CommonKVReceiver):
+26
View File
@@ -398,6 +398,32 @@ class Envs:
MOONCAKE_ENABLE_SSD_OFFLOAD = EnvBool(False) MOONCAKE_ENABLE_SSD_OFFLOAD = EnvBool(False)
MOONCAKE_OFFLOAD_FILE_STORAGE_PATH = EnvStr(None) MOONCAKE_OFFLOAD_FILE_STORAGE_PATH = EnvStr(None)
# MoRI KV Transfer
# Send CPU-resident AUX data via RDMA instead of ZMQ TCP (default: TCP).
SGLANG_MORI_SEND_AUX_RDMA = EnvBool(False)
# Number of RDMA Queue Pairs (QPs) used per transfer operation. Higher
# values can increase parallelism and bandwidth utilization.
SGLANG_MORI_QP_PER_TRANSFER = EnvInt(4)
# Number of RDMA work requests posted in a single batch to each QP. Larger
# batch sizes reduce per-operation overhead and improve throughput at the
# cost of higher latency. -1 selects automatic sizing based on the number
# of merged work requests and available endpoints.
SGLANG_MORI_POST_BATCH_SIZE = EnvInt(-1)
# Number of worker threads in the RDMA executor thread pool. More workers
# can improve parallelism for large batch transfers across multiple QPs,
# but excessive threads may cause contention.
SGLANG_MORI_NUM_WORKERS = EnvInt(4)
# Number of sharded synchronous worker threads that drain KV transfers.
# Also the bound on outstanding (posted-but-not-completed) transfers, so it
# is the primary throttle keeping the RDMA send queue from overflowing.
SGLANG_MORI_TRANSFER_SHARDS = EnvInt(8)
# Poll cadence (ms) at which a transfer worker wakes to check the SLA while
# waiting for completion; real completion still wakes it immediately.
SGLANG_MORI_WAIT_POLL_MS = EnvInt(1000)
# Per-transfer SLA (ms) before a KV transfer is failed; 0 disables the SLA
# and relies on the RDMA retry-exceeded timeout only.
SGLANG_MORI_TRANSFER_TIMEOUT_MS = EnvInt(0)
# AMD & ROCm # AMD & ROCm
SGLANG_USE_AITER = EnvBool(False) SGLANG_USE_AITER = EnvBool(False)
SGLANG_USE_AITER_AG = EnvBool(True) SGLANG_USE_AITER_AG = EnvBool(True)
@@ -8,7 +8,6 @@ from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase, PDDisaggregationServerBase,
) )
from sglang.test.test_utils import ( from sglang.test.test_utils import (
DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST,
DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
popen_launch_pd_server, popen_launch_pd_server,
@@ -179,19 +178,5 @@ class TestMoriTransferEngineTPMismatchE2E(MoriTransferEngineBase):
self._assert_generate_smoke() self._assert_generate_smoke()
class TestMoriTransferEngineHybridMambaE2E(MoriTransferEngineBase):
port_delta = 20
prefill_tp = 4
decode_tp = 4
decode_base_gpu_id = 4
required_gpus = 8
model_default = DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST
model_env_var = "SGLANG_MORI_HYBRID_E2E_TEST_MODEL"
def test_generate_smoke_hybrid_mamba(self):
self._assert_generate_smoke()
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()