[PD][MoRI] Align hybrid state transfer with per-component schema (#26539)
This commit is contained in:
@@ -35,6 +35,8 @@ from sglang.srt.disaggregation.common.conn import (
|
||||
from sglang.srt.disaggregation.common.utils import (
|
||||
AuxDataCodec,
|
||||
group_concurrent_contiguous,
|
||||
pack_int_lists,
|
||||
unpack_int_lists,
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -45,12 +47,33 @@ logger = logging.getLogger(__name__)
|
||||
MORI_GUARD = b"MoriMsgGuard"
|
||||
|
||||
|
||||
def _normalize_state_indices(
|
||||
state_indices,
|
||||
) -> Optional[npt.NDArray[np.int32]]:
|
||||
def _normalize_state_indices_per_component(
|
||||
state_indices: Optional[List],
|
||||
) -> Optional[List[Optional[npt.NDArray[np.int32]]]]:
|
||||
if state_indices is None:
|
||||
return None
|
||||
return np.asarray(state_indices, dtype=np.int32)
|
||||
out: List[Optional[npt.NDArray[np.int32]]] = []
|
||||
for entry in state_indices:
|
||||
if entry is None:
|
||||
out.append(None)
|
||||
else:
|
||||
out.append(np.asarray(entry, dtype=np.int32).ravel())
|
||||
return out
|
||||
|
||||
|
||||
def _pack_state_indices(
|
||||
state_indices: Optional[List[Optional[npt.NDArray[np.int32]]]],
|
||||
) -> bytes:
|
||||
if not state_indices:
|
||||
return b""
|
||||
lists = [(arr.tolist() if arr is not None else []) for arr in state_indices]
|
||||
return pack_int_lists(lists, "i")
|
||||
|
||||
|
||||
def _unpack_state_indices(buf: bytes) -> List[npt.NDArray[np.int32]]:
|
||||
if not buf:
|
||||
return []
|
||||
return [np.asarray(lst, dtype=np.int32) for lst in unpack_int_lists(buf, "i")]
|
||||
|
||||
|
||||
def _pack_mem_desc_list(mems: List[MemoryDesc]) -> bytes:
|
||||
@@ -67,6 +90,21 @@ def _unpack_mem_desc_list(blob: bytes) -> List[MemoryDesc]:
|
||||
return [MemoryDesc.unpack(b) for b in desc_blobs]
|
||||
|
||||
|
||||
def _pack_mem_desc_lists(mems_per_comp: List[List[MemoryDesc]]) -> bytes:
|
||||
if not mems_per_comp:
|
||||
return b""
|
||||
return msgspec.msgpack.encode(
|
||||
[[mem.pack() for mem in comp] for comp in mems_per_comp]
|
||||
)
|
||||
|
||||
|
||||
def _unpack_mem_desc_lists(blob: bytes) -> List[List[MemoryDesc]]:
|
||||
if not blob:
|
||||
return []
|
||||
nested = msgspec.msgpack.decode(blob)
|
||||
return [[MemoryDesc.unpack(b) for b in comp] for comp in nested]
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class TransferInfo:
|
||||
room: int
|
||||
@@ -75,7 +113,7 @@ class TransferInfo:
|
||||
engine_key: str
|
||||
dst_kv_indices: npt.NDArray[np.int32]
|
||||
dst_aux_index: int
|
||||
dst_state_indices: npt.NDArray[np.int32]
|
||||
dst_state_indices: List[npt.NDArray[np.int32]]
|
||||
required_dst_info_num: int
|
||||
is_dummy: bool
|
||||
|
||||
@@ -97,9 +135,9 @@ class TransferInfo:
|
||||
dst_aux_index = -1
|
||||
|
||||
if len(payload) > 6 and payload[6]:
|
||||
dst_state_indices = np.frombuffer(payload[6], dtype=np.int32)
|
||||
dst_state_indices = _unpack_state_indices(payload[6])
|
||||
else:
|
||||
dst_state_indices = np.array([], dtype=np.int32)
|
||||
dst_state_indices = []
|
||||
|
||||
required_dst_info_num = (
|
||||
int(payload[7].decode("ascii")) if len(payload) > 7 else 1
|
||||
@@ -125,13 +163,13 @@ class KVArgsRegisterInfo:
|
||||
engine_desc: EngineDesc
|
||||
dst_kv_mem_descs: List[MemoryDesc]
|
||||
dst_aux_mem_descs: List[MemoryDesc]
|
||||
dst_state_mem_descs: List[MemoryDesc]
|
||||
dst_state_mem_descs: List[List[MemoryDesc]]
|
||||
gpu_id: int
|
||||
decode_tp_size: int
|
||||
decode_tp_rank: int
|
||||
dst_kv_item_len: int
|
||||
dst_state_item_lens: List[int]
|
||||
dst_state_dim_per_tensor: List[int]
|
||||
dst_state_item_lens: List[List[int]]
|
||||
dst_state_dim_per_tensor: List[List[int]]
|
||||
|
||||
@property
|
||||
def engine_key(self) -> str:
|
||||
@@ -144,19 +182,19 @@ class KVArgsRegisterInfo:
|
||||
engine_desc = EngineDesc.unpack(payload[3])
|
||||
dst_kv_mem_descs = _unpack_mem_desc_list(payload[4])
|
||||
dst_aux_mem_descs = _unpack_mem_desc_list(payload[5])
|
||||
dst_state_mem_descs = _unpack_mem_desc_list(payload[6])
|
||||
dst_state_mem_descs = _unpack_mem_desc_lists(payload[6])
|
||||
gpu_id = int(payload[7].decode("ascii"))
|
||||
decode_tp_size = int(payload[8].decode("ascii"))
|
||||
decode_tp_rank = int(payload[9].decode("ascii"))
|
||||
dst_kv_item_len = int(payload[10].decode("ascii"))
|
||||
dst_state_item_lens = (
|
||||
list(struct.unpack(f"{len(payload[11]) // 4}I", payload[11]))
|
||||
if len(payload) > 11 and len(payload[11]) > 0
|
||||
unpack_int_lists(payload[11], "I")
|
||||
if len(payload) > 11 and payload[11]
|
||||
else []
|
||||
)
|
||||
dst_state_dim_per_tensor = (
|
||||
list(struct.unpack(f"{len(payload[12]) // 4}I", payload[12]))
|
||||
if len(payload) > 12 and len(payload[12]) > 0
|
||||
unpack_int_lists(payload[12], "I")
|
||||
if len(payload) > 12 and payload[12]
|
||||
else []
|
||||
)
|
||||
return cls(
|
||||
@@ -244,7 +282,7 @@ class MoriKVManager(CommonKVManager):
|
||||
self.engine_desc = self.engine.get_engine_desc()
|
||||
self.kv_mem_descs: List[MemoryDesc] = []
|
||||
self.aux_mem_descs: List[MemoryDesc] = []
|
||||
self.state_mem_descs: List[MemoryDesc] = []
|
||||
self.state_mem_descs: List[List[MemoryDesc]] = []
|
||||
self.transfer_lock = threading.Lock()
|
||||
self._zmq_ctx = zmq.Context()
|
||||
self._socket_local = threading.local()
|
||||
@@ -338,6 +376,7 @@ class MoriKVManager(CommonKVManager):
|
||||
self.kv_args.state_data_ptrs,
|
||||
getattr(self.kv_args, "state_data_lens", []),
|
||||
):
|
||||
component_descs: List[MemoryDesc] = []
|
||||
for ptr, length in zip(component_ptrs, component_lens):
|
||||
desc = self.engine.register_memory(
|
||||
ptr,
|
||||
@@ -345,7 +384,8 @@ class MoriKVManager(CommonKVManager):
|
||||
self.kv_args.gpu_id,
|
||||
MemoryLocationType.GPU,
|
||||
)
|
||||
self.state_mem_descs.append(desc)
|
||||
component_descs.append(desc)
|
||||
self.state_mem_descs.append(component_descs)
|
||||
|
||||
def update_status(self, bootstrap_room: int, status: KVPoll):
|
||||
current = self.request_status.get(bootstrap_room)
|
||||
@@ -904,69 +944,106 @@ class MoriKVManager(CommonKVManager):
|
||||
def send_state(
|
||||
self,
|
||||
peer_info: KVArgsRegisterInfo,
|
||||
src_state_indices: npt.NDArray[np.int32],
|
||||
dst_state_indices: npt.NDArray[np.int32],
|
||||
src_state_indices: List[npt.NDArray[np.int32]],
|
||||
dst_state_indices: List[npt.NDArray[np.int32]],
|
||||
) -> List[TransferStatus]:
|
||||
# Guard: no local state tensors -> no-op (e.g. SWA layers=0 on this PP rank)
|
||||
if not self.state_mem_descs:
|
||||
return []
|
||||
|
||||
state_type = getattr(self.kv_args, "state_type", "none")
|
||||
|
||||
if state_type == "none":
|
||||
state_types = self.kv_args.state_types
|
||||
if not state_types:
|
||||
raise RuntimeError(
|
||||
"PD state transfer failed: state_type is 'none' but state_indices were provided"
|
||||
)
|
||||
|
||||
if not peer_info.dst_state_mem_descs:
|
||||
raise RuntimeError(
|
||||
f"PD state transfer failed: remote peer has no state descriptors "
|
||||
f"(state_type={state_type}, prefill_tp_size={self.attn_tp_size}, "
|
||||
f"decode_tp_size={peer_info.decode_tp_size})"
|
||||
"PD state transfer failed: kv_args.state_types is empty but "
|
||||
"state_indices were provided"
|
||||
)
|
||||
|
||||
if len(peer_info.dst_state_mem_descs) != len(self.state_mem_descs):
|
||||
raise RuntimeError(
|
||||
f"PD state transfer failed: state descriptor count mismatch "
|
||||
f"(local={len(self.state_mem_descs)}, remote={len(peer_info.dst_state_mem_descs)}), "
|
||||
f"likely PP configuration mismatch (state_type={state_type})"
|
||||
f"PD state transfer failed: state component count mismatch "
|
||||
f"(local={len(self.state_mem_descs)}, "
|
||||
f"remote={len(peer_info.dst_state_mem_descs)})"
|
||||
)
|
||||
|
||||
if len(self.kv_args.state_item_lens) != len(self.state_mem_descs):
|
||||
raise RuntimeError(
|
||||
f"PD state transfer failed: local state_item_lens count "
|
||||
f"({len(self.kv_args.state_item_lens)}) does not match state descriptor "
|
||||
f"count ({len(self.state_mem_descs)}) (state_type={state_type})"
|
||||
src_state_item_lens = self.kv_args.state_item_lens
|
||||
src_state_dim_per_tensor = self.kv_args.state_dim_per_tensor
|
||||
|
||||
statuses: List[TransferStatus] = []
|
||||
for i, st in enumerate(state_types):
|
||||
src_indices = src_state_indices[i] if i < len(src_state_indices) else None
|
||||
dst_indices = dst_state_indices[i] if i < len(dst_state_indices) else None
|
||||
if src_indices is None or src_indices.size == 0:
|
||||
continue
|
||||
if dst_indices is None or dst_indices.size == 0:
|
||||
continue
|
||||
|
||||
src_descs = self.state_mem_descs[i]
|
||||
dst_descs = peer_info.dst_state_mem_descs[i]
|
||||
src_lens = src_state_item_lens[i] if i < len(src_state_item_lens) else []
|
||||
dst_lens = (
|
||||
peer_info.dst_state_item_lens[i]
|
||||
if i < len(peer_info.dst_state_item_lens)
|
||||
else []
|
||||
)
|
||||
src_dims = (
|
||||
src_state_dim_per_tensor[i] if i < len(src_state_dim_per_tensor) else []
|
||||
)
|
||||
dst_dims = (
|
||||
peer_info.dst_state_dim_per_tensor[i]
|
||||
if i < len(peer_info.dst_state_dim_per_tensor)
|
||||
else []
|
||||
)
|
||||
|
||||
if state_type == "mamba":
|
||||
return self._send_mamba_state(
|
||||
peer_info, src_state_indices, dst_state_indices
|
||||
)
|
||||
elif state_type in ("swa", "dsa"):
|
||||
return self._send_swa_dsa_state(
|
||||
peer_info, src_state_indices, dst_state_indices, state_type
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"PD state transfer failed: unknown state_type={state_type}"
|
||||
)
|
||||
if st == "mamba":
|
||||
statuses.extend(
|
||||
self._send_mamba_state(
|
||||
peer_info,
|
||||
src_indices,
|
||||
dst_indices,
|
||||
src_descs,
|
||||
dst_descs,
|
||||
src_lens,
|
||||
dst_lens,
|
||||
src_dims,
|
||||
dst_dims,
|
||||
)
|
||||
)
|
||||
elif st in ("swa", "dsa"):
|
||||
statuses.extend(
|
||||
self._send_swa_dsa_state(
|
||||
peer_info,
|
||||
src_indices,
|
||||
dst_indices,
|
||||
src_descs,
|
||||
src_lens,
|
||||
dst_descs,
|
||||
st,
|
||||
)
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(f"PD state transfer failed: unknown state_type={st}")
|
||||
|
||||
return statuses
|
||||
|
||||
def _send_mamba_state(
|
||||
self,
|
||||
peer_info: KVArgsRegisterInfo,
|
||||
src_state_indices: npt.NDArray[np.int32],
|
||||
dst_state_indices: npt.NDArray[np.int32],
|
||||
src_state_mem_descs: List[MemoryDesc],
|
||||
dst_state_mem_descs: List[MemoryDesc],
|
||||
src_state_item_lens: List[int],
|
||||
dst_state_item_lens: List[int],
|
||||
src_state_dim_per_tensor: List[int],
|
||||
dst_state_dim_per_tensor: List[int],
|
||||
) -> List[TransferStatus]:
|
||||
if len(src_state_indices) != 1 or len(dst_state_indices) != 1:
|
||||
if src_state_indices.size != 1 or dst_state_indices.size != 1:
|
||||
raise RuntimeError(
|
||||
f"PD state transfer failed: mamba requires single state index, "
|
||||
f"got src={len(src_state_indices)}, dst={len(dst_state_indices)}"
|
||||
f"got src={src_state_indices.size}, dst={dst_state_indices.size}"
|
||||
)
|
||||
|
||||
tp_mismatch = peer_info.decode_tp_size != self.attn_tp_size
|
||||
src_state_dim_per_tensor = getattr(self.kv_args, "state_dim_per_tensor", [])
|
||||
dst_state_dim_per_tensor = peer_info.dst_state_dim_per_tensor
|
||||
|
||||
# If dim info missing, silently degrade to whole-item copy (Mooncake compat)
|
||||
if tp_mismatch and (
|
||||
@@ -983,15 +1060,14 @@ class MoriKVManager(CommonKVManager):
|
||||
|
||||
src_idx = int(src_state_indices[0])
|
||||
dst_idx = int(dst_state_indices[0])
|
||||
statuses = []
|
||||
statuses: List[TransferStatus] = []
|
||||
|
||||
local_tp_rank = self.kv_args.engine_rank % self.attn_tp_size
|
||||
dst_tp_rank = peer_info.decode_tp_rank % peer_info.decode_tp_size
|
||||
|
||||
for i in range(len(self.state_mem_descs)):
|
||||
src_desc = self.state_mem_descs[i]
|
||||
dst_desc = peer_info.dst_state_mem_descs[i]
|
||||
src_item_len = self.kv_args.state_item_lens[i]
|
||||
for i, src_desc in enumerate(src_state_mem_descs):
|
||||
dst_desc = dst_state_mem_descs[i]
|
||||
src_item_len = src_state_item_lens[i]
|
||||
|
||||
if not tp_mismatch:
|
||||
# same-TP: whole item copy
|
||||
@@ -1000,7 +1076,7 @@ class MoriKVManager(CommonKVManager):
|
||||
size = src_item_len
|
||||
else:
|
||||
# TP mismatch slice copy
|
||||
dst_item_len = peer_info.dst_state_item_lens[i]
|
||||
dst_item_len = dst_state_item_lens[i]
|
||||
src_dim = src_state_dim_per_tensor[i]
|
||||
dst_dim = dst_state_dim_per_tensor[i]
|
||||
|
||||
@@ -1044,6 +1120,9 @@ class MoriKVManager(CommonKVManager):
|
||||
peer_info: KVArgsRegisterInfo,
|
||||
src_state_indices: npt.NDArray[np.int32],
|
||||
dst_state_indices: npt.NDArray[np.int32],
|
||||
src_state_mem_descs: List[MemoryDesc],
|
||||
src_state_item_lens: List[int],
|
||||
dst_state_mem_descs: List[MemoryDesc],
|
||||
state_type: str,
|
||||
) -> List[TransferStatus]:
|
||||
# TP mismatch check for non-MLA SWA
|
||||
@@ -1057,17 +1136,17 @@ class MoriKVManager(CommonKVManager):
|
||||
f"(prefill_tp_size={self.attn_tp_size}, decode_tp_size={peer_info.decode_tp_size})"
|
||||
)
|
||||
|
||||
common_len = min(len(src_state_indices), len(dst_state_indices))
|
||||
if common_len == 0 and max(len(src_state_indices), len(dst_state_indices)) > 0:
|
||||
common_len = min(src_state_indices.size, dst_state_indices.size)
|
||||
if common_len == 0 and max(src_state_indices.size, dst_state_indices.size) > 0:
|
||||
raise RuntimeError(
|
||||
f"No overlapping state indices for state_type={state_type}"
|
||||
)
|
||||
if len(src_state_indices) != len(dst_state_indices):
|
||||
if src_state_indices.size != dst_state_indices.size:
|
||||
logger.warning(
|
||||
"State index length mismatch for %s: src=%d dst=%d; truncating to common prefix=%d",
|
||||
state_type,
|
||||
len(src_state_indices),
|
||||
len(dst_state_indices),
|
||||
src_state_indices.size,
|
||||
dst_state_indices.size,
|
||||
common_len,
|
||||
)
|
||||
src_state_indices = src_state_indices[:common_len]
|
||||
@@ -1078,11 +1157,10 @@ class MoriKVManager(CommonKVManager):
|
||||
*group_concurrent_contiguous(src_state_indices, dst_state_indices)
|
||||
)
|
||||
|
||||
statuses = []
|
||||
for i in range(len(self.state_mem_descs)):
|
||||
src_desc = self.state_mem_descs[i]
|
||||
dst_desc = peer_info.dst_state_mem_descs[i]
|
||||
state_item_len = self.kv_args.state_item_lens[i]
|
||||
statuses: List[TransferStatus] = []
|
||||
for i, src_desc in enumerate(src_state_mem_descs):
|
||||
dst_desc = dst_state_mem_descs[i]
|
||||
state_item_len = src_state_item_lens[i]
|
||||
|
||||
statuses.extend(
|
||||
self._submit_batch_transfer_plan(
|
||||
@@ -1117,7 +1195,7 @@ class MoriKVManager(CommonKVManager):
|
||||
index_slice: slice,
|
||||
is_last_chunk: bool,
|
||||
aux_index: Optional[int] = None,
|
||||
state_indices: Optional[npt.NDArray[np.int32]] = None,
|
||||
state_indices: Optional[List[npt.NDArray[np.int32]]] = None,
|
||||
) -> Tuple[List[TransferStatus], Optional[List[TransferInfo]]]:
|
||||
assert self.disaggregation_mode == DisaggregationMode.PREFILL
|
||||
|
||||
@@ -1233,7 +1311,9 @@ class MoriKVSender(CommonKVSender):
|
||||
return
|
||||
|
||||
normalized_state = (
|
||||
_normalize_state_indices(state_indices) if is_last_chunk else None
|
||||
_normalize_state_indices_per_component(state_indices)
|
||||
if is_last_chunk
|
||||
else None
|
||||
)
|
||||
statuses, infos = self.kv_mgr.add_transfer_request(
|
||||
self.bootstrap_room,
|
||||
@@ -1382,18 +1462,16 @@ class MoriKVReceiver(CommonKVReceiver):
|
||||
engine_desc_blob = self.kv_mgr.engine_desc.pack()
|
||||
packed_kv_descs = _pack_mem_desc_list(self.kv_mgr.kv_mem_descs)
|
||||
packed_aux_descs = _pack_mem_desc_list(self.kv_mgr.aux_mem_descs)
|
||||
packed_state_descs = _pack_mem_desc_list(self.kv_mgr.state_mem_descs)
|
||||
packed_state_descs = _pack_mem_desc_lists(self.kv_mgr.state_mem_descs)
|
||||
gpu_id = str(self.kv_mgr.kv_args.gpu_id).encode("ascii")
|
||||
decode_tp_size = str(self.kv_mgr.attn_tp_size).encode("ascii")
|
||||
decode_tp_rank = str(self.kv_mgr.kv_args.engine_rank).encode("ascii")
|
||||
kv_item_len = str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii")
|
||||
packed_state_item_lens = b"".join(
|
||||
struct.pack("I", item_len)
|
||||
for item_len in self.kv_mgr.kv_args.state_item_lens
|
||||
packed_state_item_lens = pack_int_lists(
|
||||
self.kv_mgr.kv_args.state_item_lens, "I"
|
||||
)
|
||||
state_dim_per_tensor = getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", [])
|
||||
packed_state_dim_per_tensor = b"".join(
|
||||
struct.pack("I", dim) for dim in state_dim_per_tensor
|
||||
packed_state_dim_per_tensor = pack_int_lists(
|
||||
self.kv_mgr.kv_args.state_dim_per_tensor, "I"
|
||||
)
|
||||
|
||||
for bootstrap_info in self.bootstrap_infos:
|
||||
@@ -1432,13 +1510,13 @@ class MoriKVReceiver(CommonKVReceiver):
|
||||
np.asarray(kv_indices, dtype=np.int32).tobytes() if kv_indices.size else b""
|
||||
)
|
||||
aux_bytes = str(aux_index).encode("ascii") if aux_index is not None else b""
|
||||
normalized_state = _normalize_state_indices(state_indices)
|
||||
normalized_state = _normalize_state_indices_per_component(state_indices)
|
||||
|
||||
for bootstrap_info in self.bootstrap_infos:
|
||||
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
|
||||
is_dummy = bootstrap_info.get("is_dummy", False)
|
||||
if not is_dummy and normalized_state is not None:
|
||||
state_bytes = normalized_state.tobytes()
|
||||
state_bytes = _pack_state_indices(normalized_state)
|
||||
else:
|
||||
state_bytes = b""
|
||||
with lock:
|
||||
|
||||
@@ -8,13 +8,14 @@ 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,
|
||||
try_cached_model,
|
||||
)
|
||||
|
||||
register_amd_ci(est_time=300, suite="stage-b-test-large-8-gpu-mi35x-disaggregation-amd")
|
||||
register_amd_ci(est_time=900, suite="stage-b-test-large-8-gpu-mi35x-disaggregation-amd")
|
||||
|
||||
|
||||
class MoriTransferEngineBase(PDDisaggregationServerBase):
|
||||
@@ -24,6 +25,12 @@ class MoriTransferEngineBase(PDDisaggregationServerBase):
|
||||
decode_base_gpu_id = 1
|
||||
required_gpus = 2
|
||||
|
||||
# Subclasses can override to pick a different model or pass extra args.
|
||||
model_default = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||
model_env_var = "SGLANG_MORI_E2E_TEST_MODEL"
|
||||
extra_prefill_args: list = []
|
||||
extra_decode_args: list = []
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
try:
|
||||
@@ -56,10 +63,7 @@ class MoriTransferEngineBase(PDDisaggregationServerBase):
|
||||
|
||||
cls._shift_ports()
|
||||
cls.model = try_cached_model(
|
||||
os.environ.get(
|
||||
"SGLANG_MORI_E2E_TEST_MODEL",
|
||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||
)
|
||||
os.environ.get(cls.model_env_var, cls.model_default)
|
||||
)
|
||||
|
||||
cls.start_prefill()
|
||||
@@ -111,7 +115,7 @@ class MoriTransferEngineBase(PDDisaggregationServerBase):
|
||||
str(cls.prefill_tp),
|
||||
"--attention-backend",
|
||||
"aiter",
|
||||
]
|
||||
] + list(cls.extra_prefill_args)
|
||||
prefill_args += cls.transfer_backend + cls.rdma_devices
|
||||
cls.process_prefill = popen_launch_pd_server(
|
||||
cls.model,
|
||||
@@ -134,7 +138,7 @@ class MoriTransferEngineBase(PDDisaggregationServerBase):
|
||||
str(cls.decode_base_gpu_id),
|
||||
"--attention-backend",
|
||||
"aiter",
|
||||
]
|
||||
] + list(cls.extra_decode_args)
|
||||
decode_args += cls.transfer_backend + cls.rdma_devices
|
||||
cls.process_decode = popen_launch_pd_server(
|
||||
cls.model,
|
||||
@@ -175,5 +179,19 @@ 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()
|
||||
|
||||
Reference in New Issue
Block a user