From 18d728967a72c285d3ef98b339f63550b0f68b7b Mon Sep 17 00:00:00 2001 From: Niko Ma Date: Mon, 8 Jun 2026 15:49:25 +0800 Subject: [PATCH] [PD][MoRI] Drive KV transfers with a sharded synchronous worker pool (#26922) --- docker/rocm.Dockerfile | 2 +- python/sglang/srt/disaggregation/mori/conn.py | 304 ++++++++++++------ python/sglang/srt/environ.py | 26 ++ .../test_mori_transfer_engine_e2e.py | 15 - 4 files changed, 235 insertions(+), 112 deletions(-) diff --git a/docker/rocm.Dockerfile b/docker/rocm.Dockerfile index ef8f2e910..d1f3afcae 100644 --- a/docker/rocm.Dockerfile +++ b/docker/rocm.Dockerfile @@ -104,7 +104,7 @@ ARG ENABLE_MORI=0 ARG NIC_BACKEND=none ARG MORI_REPO="https://github.com/ROCm/mori.git" -ARG MORI_COMMIT="96ffa169710f214e76e07abe5008d686fe54522b" +ARG MORI_COMMIT="d87651c998296c5bddbeefc4cc525b58663b1636" # AMD AINIC apt repo settings ARG AINIC_VERSION=1.117.5-a-38 diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index 4c392e466..c1e6ae22f 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -23,6 +23,7 @@ from mori.io import ( MemoryLocationType, PollCqMode, RdmaBackendConfig, + StatusCode, ) 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 ( AuxDataCodec, + FastQueue, group_concurrent_contiguous, pack_int_lists, unpack_int_lists, ) from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.environ import envs 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 logger = logging.getLogger(__name__) @@ -267,6 +269,16 @@ class TransferTarget: 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): AUX_DATA_HEADER = b"AUX_DATA" @@ -286,13 +298,25 @@ class MoriKVManager(CommonKVManager): self.transfer_lock = threading.Lock() self._zmq_ctx = zmq.Context() self._socket_local = threading.local() - # Send CPU-resident AUX data via RDMA instead of ZMQ TCP. - # 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._send_aux_rdma = envs.SGLANG_MORI_SEND_AUX_RDMA.get() self._register_local_buffers() 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() elif self.disaggregation_mode == DisaggregationMode.DECODE: self.room_to_bootstrap_addr: Dict[int, str] = {} @@ -315,24 +339,9 @@ class MoriKVManager(CommonKVManager): engine = IOEngine(engine_key, config) poll_mode = PollCqMode.POLLING - # Number of RDMA Queue Pairs (QPs) used per transfer operation. - # Higher values can increase parallelism and bandwidth utilization. - # Default: 4 - 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) + qp_per_transfer = envs.SGLANG_MORI_QP_PER_TRANSFER.get() + post_batch_size = envs.SGLANG_MORI_POST_BATCH_SIZE.get() + num_worker_threads = envs.SGLANG_MORI_NUM_WORKERS.get() rdma_cfg = RdmaBackendConfig( qp_per_transfer, @@ -400,6 +409,34 @@ class MoriKVManager(CommonKVManager): return 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): """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] 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 = ( self._get_mla_mem_desc_slices(peer_info.dst_kv_mem_descs) ) 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( self._submit_batch_transfer_plan( src_descs[layer_id], @@ -1273,12 +1312,6 @@ class MoriKVManager(CommonKVManager): ) 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 @@ -1294,10 +1327,12 @@ class MoriKVSender(CommonKVSender): super().__init__(mgr, bootstrap_addr, bootstrap_room, dest_tp_ranks, pp_rank) self.transfer_statuses: List[TransferStatus] = [] self.pending_infos: Optional[List[TransferInfo]] = None - self.sent_last_chunk = False self.conclude_state: Optional[KVPoll] = None self.status_notified = False self.init_time = time.time() + self._notify_lock = threading.Lock() + self._notified_status: Optional[KVPoll] = None + self._notified_reason: Optional[str] = None def send( self, @@ -1315,20 +1350,17 @@ class MoriKVSender(CommonKVSender): if is_last_chunk else None ) - statuses, infos = self.kv_mgr.add_transfer_request( - self.bootstrap_room, - kv_indices, - index_slice, - is_last_chunk, - aux_index=self.aux_index if is_last_chunk else None, - state_indices=normalized_state, + self._record_transfer_indices(kv_indices, state_indices) + self.kv_mgr.enqueue_transfer( + _TransferChunk( + sender=self, + kv_indices=kv_indices, + index_slice=index_slice, + 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() 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: 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: if self.conclude_state is not None: return self.conclude_state if self.bootstrap_room not in self.kv_mgr.request_status: - self._finalize_failure() - return KVPoll.Failed + sent_status, _ = self._finalize_failure() + return sent_status status = self.kv_mgr.check_status(self.bootstrap_room) if status == KVPoll.Bootstrapping: - timeout_result = self._check_bootstrap_timeout() - if timeout_result is not None: - self._finalize_failure() - return KVPoll.Failed + elapsed = time.time() - self.init_time + if elapsed >= self.kv_mgr.bootstrap_timeout: + logger.warning_once( + "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 if status == KVPoll.Failed: - self._finalize_failure() - return KVPoll.Failed + sent_status, _ = self._finalize_failure() + return sent_status - if status == KVPoll.Success and self.kv_mgr.is_dummy_cp_rank: + if status == KVPoll.Success: self.conclude_state = KVPoll.Success return KVPoll.Success - transfers_done = self._all_transfers_finished() - 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) + return status def _collect_failure_reason(self) -> str: for status in self.transfer_statuses: @@ -1391,33 +1473,66 @@ class MoriKVSender(CommonKVSender): return f"KV transfer failed: {status.Message()}" return "KV transfer failed due to unknown reason" - def _notify_decode( - self, status: KVPoll, failure_reason: Optional[str] = None - ) -> None: + def _terminalize_locked( + self, + status: KVPoll, + reason: Optional[str] = None, + ) -> Tuple[KVPoll, Optional[str], Optional[List[TransferInfo]]]: 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 if infos is None: with self.kv_mgr.transfer_lock: room_infos = self.kv_mgr.transfer_infos.get(self.bootstrap_room) - if room_infos is not None: - infos = list(room_infos.values()) + infos = list(room_infos.values()) if room_infos is not None else None + + 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: 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: - if self.conclude_state == KVPoll.Failed: - return + def _finalize_failure( + self, failure_reason: Optional[str] = None + ) -> Tuple[KVPoll, Optional[str]]: if failure_reason is None: - 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 + with self.kv_mgr.failure_lock: + failure_reason = self.kv_mgr.failure_records.get( + self.bootstrap_room, "KV transfer 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): if self.conclude_state is None: @@ -1430,10 +1545,7 @@ class MoriKVSender(CommonKVSender): raise RuntimeError(failure_reason) def abort(self): - self.kv_mgr.record_failure(self.bootstrap_room, "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 + self._finalize_failure("Aborted by AbortReq.") class MoriKVReceiver(CommonKVReceiver): diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 4c6684c74..5a4323264 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -398,6 +398,32 @@ class Envs: MOONCAKE_ENABLE_SSD_OFFLOAD = EnvBool(False) 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 SGLANG_USE_AITER = EnvBool(False) SGLANG_USE_AITER_AG = EnvBool(True) diff --git a/test/registered/amd/disaggregation/test_mori_transfer_engine_e2e.py b/test/registered/amd/disaggregation/test_mori_transfer_engine_e2e.py index 9fa13413e..1adebc652 100644 --- a/test/registered/amd/disaggregation/test_mori_transfer_engine_e2e.py +++ b/test/registered/amd/disaggregation/test_mori_transfer_engine_e2e.py @@ -8,7 +8,6 @@ from sglang.test.server_fixtures.disaggregation_fixture import ( PDDisaggregationServerBase, ) from sglang.test.test_utils import ( - DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, popen_launch_pd_server, @@ -179,19 +178,5 @@ class TestMoriTransferEngineTPMismatchE2E(MoriTransferEngineBase): 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__": unittest.main()