Support PD disaggregation with different TP/DP size for Qwen3-Next (#16056)
Co-authored-by: xjx392321 <xjx392321@alibaba-inc.com>
This commit is contained in:
co-authored by
xjx392321
parent
30ece5e1d6
commit
9bd92ba0f6
@@ -24,6 +24,8 @@ class KVArgs:
|
|||||||
state_data_lens: List[int]
|
state_data_lens: List[int]
|
||||||
state_item_lens: List[int]
|
state_item_lens: List[int]
|
||||||
state_type: str # "none", "mamba", "swa"
|
state_type: str # "none", "mamba", "swa"
|
||||||
|
# for mamba state different tp slice transfer
|
||||||
|
state_dim_per_tensor: List[int] # dimension to slice for each state tensor
|
||||||
ib_device: str
|
ib_device: str
|
||||||
ib_traffic_class: str
|
ib_traffic_class: str
|
||||||
gpu_id: int
|
gpu_id: int
|
||||||
|
|||||||
@@ -292,6 +292,11 @@ class DecodePreallocQueue:
|
|||||||
kv_args.state_type = "swa"
|
kv_args.state_type = "swa"
|
||||||
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
|
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
|
||||||
kv_args.state_type = "mamba"
|
kv_args.state_type = "mamba"
|
||||||
|
# Get state dimension info for cross-TP slice transfer
|
||||||
|
if hasattr(self.token_to_kv_pool, "get_state_dim_per_tensor"):
|
||||||
|
kv_args.state_dim_per_tensor = (
|
||||||
|
self.token_to_kv_pool.get_state_dim_per_tensor()
|
||||||
|
)
|
||||||
elif isinstance(self.token_to_kv_pool, NSATokenToKVPool):
|
elif isinstance(self.token_to_kv_pool, NSATokenToKVPool):
|
||||||
kv_args.state_type = "nsa"
|
kv_args.state_type = "nsa"
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -114,6 +114,9 @@ class KVArgsRegisterInfo:
|
|||||||
dst_tp_rank: int
|
dst_tp_rank: int
|
||||||
dst_attn_tp_size: int
|
dst_attn_tp_size: int
|
||||||
dst_kv_item_len: int
|
dst_kv_item_len: int
|
||||||
|
# for mamba state different tp slice transfer
|
||||||
|
dst_state_item_lens: list[int]
|
||||||
|
dst_state_dim_per_tensor: list[int]
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_zmq(cls, msg: List[bytes]):
|
def from_zmq(cls, msg: List[bytes]):
|
||||||
@@ -128,6 +131,16 @@ class KVArgsRegisterInfo:
|
|||||||
dst_tp_rank=int(msg[7].decode("ascii")),
|
dst_tp_rank=int(msg[7].decode("ascii")),
|
||||||
dst_attn_tp_size=int(msg[8].decode("ascii")),
|
dst_attn_tp_size=int(msg[8].decode("ascii")),
|
||||||
dst_kv_item_len=int(msg[9].decode("ascii")),
|
dst_kv_item_len=int(msg[9].decode("ascii")),
|
||||||
|
dst_state_item_lens=(
|
||||||
|
list(struct.unpack(f"{len(msg[10])//4}I", msg[10]))
|
||||||
|
if len(msg) > 10 and len(msg[10]) > 0
|
||||||
|
else []
|
||||||
|
),
|
||||||
|
dst_state_dim_per_tensor=(
|
||||||
|
list(struct.unpack(f"{len(msg[11])//4}I", msg[11]))
|
||||||
|
if len(msg) > 11 and len(msg[11]) > 0
|
||||||
|
else []
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -641,17 +654,42 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
prefill_state_indices: list[int],
|
prefill_state_indices: list[int],
|
||||||
dst_state_data_ptrs: list[int],
|
dst_state_data_ptrs: list[int],
|
||||||
executor: concurrent.futures.ThreadPoolExecutor,
|
executor: concurrent.futures.ThreadPoolExecutor,
|
||||||
|
target_rank_registration_info: Optional[KVArgsRegisterInfo] = None,
|
||||||
):
|
):
|
||||||
"""Send state or extra pool data with type-specific handling."""
|
"""Send state or extra pool data with type-specific handling."""
|
||||||
state_type = getattr(self.kv_args, "state_type", "none")
|
state_type = getattr(self.kv_args, "state_type", "none")
|
||||||
|
|
||||||
if state_type == "mamba":
|
if state_type == "mamba":
|
||||||
return self._send_mamba_state(
|
# Check if we need slice transfer for different TP sizes
|
||||||
req,
|
if (
|
||||||
prefill_state_indices,
|
target_rank_registration_info is not None
|
||||||
dst_state_data_ptrs,
|
and self.attn_tp_size != target_rank_registration_info.dst_attn_tp_size
|
||||||
)
|
):
|
||||||
|
return self._send_mamba_state_slice(
|
||||||
|
req,
|
||||||
|
prefill_state_indices,
|
||||||
|
dst_state_data_ptrs,
|
||||||
|
target_rank_registration_info.dst_state_item_lens,
|
||||||
|
target_rank_registration_info.dst_state_dim_per_tensor,
|
||||||
|
target_rank_registration_info.dst_tp_rank,
|
||||||
|
target_rank_registration_info.dst_attn_tp_size,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
return self._send_mamba_state(
|
||||||
|
req,
|
||||||
|
prefill_state_indices,
|
||||||
|
dst_state_data_ptrs,
|
||||||
|
)
|
||||||
elif state_type in ["swa", "nsa"]:
|
elif state_type in ["swa", "nsa"]:
|
||||||
|
# SWA and NSA hybrid models do not support different TP sizes yet
|
||||||
|
if (
|
||||||
|
target_rank_registration_info is not None
|
||||||
|
and not self.is_mla_backend
|
||||||
|
and self.attn_tp_size != target_rank_registration_info.dst_attn_tp_size
|
||||||
|
):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet."
|
||||||
|
)
|
||||||
# Reuse _send_kvcache_generic interface to send extra pool data
|
# Reuse _send_kvcache_generic interface to send extra pool data
|
||||||
prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32)
|
prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32)
|
||||||
dst_state_indices = np.array(req.dst_state_indices, dtype=np.int32)
|
dst_state_indices = np.array(req.dst_state_indices, dtype=np.int32)
|
||||||
@@ -688,6 +726,90 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
|
|
||||||
return self._transfer_data(req.mooncake_session_id, transfer_blocks)
|
return self._transfer_data(req.mooncake_session_id, transfer_blocks)
|
||||||
|
|
||||||
|
def _send_mamba_state_slice(
|
||||||
|
self,
|
||||||
|
req: TransferInfo,
|
||||||
|
prefill_mamba_index: list[int],
|
||||||
|
dst_state_data_ptrs: list[int],
|
||||||
|
dst_state_item_lens: list[int],
|
||||||
|
dst_state_dim_per_tensor: list[int],
|
||||||
|
dst_tp_rank: int,
|
||||||
|
dst_attn_tp_size: int,
|
||||||
|
):
|
||||||
|
"""Transfer Mamba states with TP slice support.
|
||||||
|
|
||||||
|
Mamba state layout:
|
||||||
|
- conv_state: [num_layers, size+1, conv_dim/tp, conv_kernel-1]
|
||||||
|
- temporal_state: [num_layers, size+1, num_heads/tp, head_dim, state_size]
|
||||||
|
|
||||||
|
The 3rd dimension is sliced by TP. When prefill and decode have different
|
||||||
|
attn_tp_size, we need to slice the state accordingly.
|
||||||
|
"""
|
||||||
|
logger.warning_once(
|
||||||
|
"Using Mamba state slice transfer for different TP sizes between prefill and decode. "
|
||||||
|
f"Prefill attn_tp_size={self.attn_tp_size}, Decode attn_tp_size={dst_attn_tp_size}. "
|
||||||
|
"Performance may be affected."
|
||||||
|
)
|
||||||
|
assert len(prefill_mamba_index) == 1, "Mamba should have single state index"
|
||||||
|
|
||||||
|
transfer_blocks = []
|
||||||
|
prefill_state_data_ptrs = self.kv_args.state_data_ptrs
|
||||||
|
prefill_state_item_lens = self.kv_args.state_item_lens
|
||||||
|
src_state_dim_per_tensor = getattr(self.kv_args, "state_dim_per_tensor", [])
|
||||||
|
|
||||||
|
# If no dimension info available, fall back to regular transfer
|
||||||
|
if not src_state_dim_per_tensor or not dst_state_dim_per_tensor:
|
||||||
|
return self._send_mamba_state(req, prefill_mamba_index, dst_state_data_ptrs)
|
||||||
|
|
||||||
|
local_tp_rank_in_group = self.kv_args.engine_rank % self.attn_tp_size
|
||||||
|
dst_tp_rank_in_group = dst_tp_rank % dst_attn_tp_size
|
||||||
|
|
||||||
|
for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
|
||||||
|
src_item_len = prefill_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]
|
||||||
|
|
||||||
|
# Calculate bytes per dimension slice
|
||||||
|
# item_len = dim * trailing_dims_size, so trailing_dims_size = item_len / dim
|
||||||
|
src_bytes_per_dim = src_item_len // src_dim
|
||||||
|
dst_bytes_per_dim = dst_item_len // dst_dim
|
||||||
|
|
||||||
|
# Determine slicing parameters based on TP configuration
|
||||||
|
if self.attn_tp_size > dst_attn_tp_size:
|
||||||
|
# Multiple prefill ranks send to 1 decode rank
|
||||||
|
# Each prefill sends all its dims to the appropriate offset in decode
|
||||||
|
src_dim_start = 0
|
||||||
|
num_dims_to_send = src_dim
|
||||||
|
dst_dim_start = local_tp_rank_in_group * src_dim
|
||||||
|
else:
|
||||||
|
# 1 prefill rank sends to multiple decode ranks
|
||||||
|
# Prefill sends a slice of its dims to each decode rank
|
||||||
|
src_dim_start = (dst_tp_rank_in_group * dst_dim) % src_dim
|
||||||
|
num_dims_to_send = dst_dim
|
||||||
|
dst_dim_start = 0
|
||||||
|
|
||||||
|
# Calculate byte offsets
|
||||||
|
src_dim_offset = src_dim_start * src_bytes_per_dim
|
||||||
|
dst_dim_offset = dst_dim_start * dst_bytes_per_dim
|
||||||
|
bytes_to_send = num_dims_to_send * src_bytes_per_dim
|
||||||
|
|
||||||
|
# Calculate addresses for this state tensor
|
||||||
|
src_addr = (
|
||||||
|
prefill_state_data_ptrs[i]
|
||||||
|
+ src_item_len * int(prefill_mamba_index[0])
|
||||||
|
+ src_dim_offset
|
||||||
|
)
|
||||||
|
dst_addr = (
|
||||||
|
dst_state_ptr
|
||||||
|
+ dst_item_len * int(req.dst_state_indices[0])
|
||||||
|
+ dst_dim_offset
|
||||||
|
)
|
||||||
|
|
||||||
|
transfer_blocks.append((src_addr, dst_addr, bytes_to_send))
|
||||||
|
|
||||||
|
return self._transfer_data(req.mooncake_session_id, transfer_blocks)
|
||||||
|
|
||||||
def sync_status_to_decode_endpoint(
|
def sync_status_to_decode_endpoint(
|
||||||
self, remote: str, dst_port: int, room: int, status: int, prefill_rank: int
|
self, remote: str, dst_port: int, room: int, status: int, prefill_rank: int
|
||||||
):
|
):
|
||||||
@@ -798,19 +920,12 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
|
|
||||||
if kv_chunk.is_last:
|
if kv_chunk.is_last:
|
||||||
if kv_chunk.state_indices is not None:
|
if kv_chunk.state_indices is not None:
|
||||||
if not self.is_mla_backend and (
|
|
||||||
self.attn_tp_size
|
|
||||||
!= target_rank_registration_info.dst_attn_tp_size
|
|
||||||
):
|
|
||||||
raise RuntimeError(
|
|
||||||
f"PD Disaggregation does NOT support PD different TP sizes for non-MLA hybrid models yet."
|
|
||||||
)
|
|
||||||
|
|
||||||
self.maybe_send_extra(
|
self.maybe_send_extra(
|
||||||
req,
|
req,
|
||||||
kv_chunk.state_indices,
|
kv_chunk.state_indices,
|
||||||
target_rank_registration_info.dst_state_data_ptrs,
|
target_rank_registration_info.dst_state_data_ptrs,
|
||||||
executor,
|
executor,
|
||||||
|
target_rank_registration_info,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Only the last chunk we need to send the aux data
|
# Only the last chunk we need to send the aux data
|
||||||
@@ -1204,6 +1319,17 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
|||||||
packed_state_data_ptrs = b"".join(
|
packed_state_data_ptrs = b"".join(
|
||||||
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs
|
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs
|
||||||
)
|
)
|
||||||
|
# Pack state_item_lens and state_dim_per_tensor for mamba state slice transfer
|
||||||
|
packed_state_item_lens = b"".join(
|
||||||
|
struct.pack("I", item_len)
|
||||||
|
for item_len in self.kv_mgr.kv_args.state_item_lens
|
||||||
|
)
|
||||||
|
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
|
||||||
|
)
|
||||||
# Note(shangming): No need to add pp rank here since decode pp size should be equal to prefill pp size or 1
|
# Note(shangming): No need to add pp rank here since decode pp size should be equal to prefill pp size or 1
|
||||||
tp_rank = self.kv_mgr.kv_args.engine_rank
|
tp_rank = self.kv_mgr.kv_args.engine_rank
|
||||||
kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0]
|
kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0]
|
||||||
@@ -1225,6 +1351,8 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
|||||||
dst_tp_rank,
|
dst_tp_rank,
|
||||||
dst_attn_tp_size,
|
dst_attn_tp_size,
|
||||||
dst_kv_item_len,
|
dst_kv_item_len,
|
||||||
|
packed_state_item_lens,
|
||||||
|
packed_state_dim_per_tensor,
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -161,6 +161,11 @@ class PrefillBootstrapQueue:
|
|||||||
kv_args.state_type = "swa"
|
kv_args.state_type = "swa"
|
||||||
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
|
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
|
||||||
kv_args.state_type = "mamba"
|
kv_args.state_type = "mamba"
|
||||||
|
# Get state dimension info for cross-TP slice transfer
|
||||||
|
if hasattr(self.token_to_kv_pool, "get_state_dim_per_tensor"):
|
||||||
|
kv_args.state_dim_per_tensor = (
|
||||||
|
self.token_to_kv_pool.get_state_dim_per_tensor()
|
||||||
|
)
|
||||||
elif isinstance(self.token_to_kv_pool, NSATokenToKVPool):
|
elif isinstance(self.token_to_kv_pool, NSATokenToKVPool):
|
||||||
kv_args.state_type = "nsa"
|
kv_args.state_type = "nsa"
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -374,6 +374,33 @@ class MambaPool:
|
|||||||
]
|
]
|
||||||
return data_ptrs, data_lens, item_lens
|
return data_ptrs, data_lens, item_lens
|
||||||
|
|
||||||
|
def get_state_dim_per_tensor(self):
|
||||||
|
"""Get the sliceable dimension size for each state tensor.
|
||||||
|
|
||||||
|
For mamba state, the layout is:
|
||||||
|
- conv_state: [num_layers, size+1, conv_dim/tp, conv_kernel-1]
|
||||||
|
- temporal_state: [num_layers, size+1, num_heads/tp, head_dim, state_size]
|
||||||
|
|
||||||
|
The 3rd dimension (index 2) is the one that gets sliced by TP.
|
||||||
|
Returns the size of this dimension for each tensor (repeated for each layer).
|
||||||
|
"""
|
||||||
|
state_tensors = []
|
||||||
|
for field in vars(self.mamba_cache):
|
||||||
|
value = getattr(self.mamba_cache, field)
|
||||||
|
if isinstance(value, list):
|
||||||
|
state_tensors.extend(value)
|
||||||
|
else:
|
||||||
|
state_tensors.append(value)
|
||||||
|
|
||||||
|
dim_per_tensor = []
|
||||||
|
for state_tensor in state_tensors:
|
||||||
|
# state_tensor shape: [num_layers, size+1, sliceable_dim, ...]
|
||||||
|
# The sliceable dimension is at index 2 (after num_layers and size)
|
||||||
|
sliceable_dim = state_tensor.shape[2]
|
||||||
|
# Repeat for each layer since we have per-layer data_ptrs
|
||||||
|
dim_per_tensor += [sliceable_dim] * self.num_mamba_layers
|
||||||
|
return dim_per_tensor
|
||||||
|
|
||||||
|
|
||||||
class HybridReqToTokenPool(ReqToTokenPool):
|
class HybridReqToTokenPool(ReqToTokenPool):
|
||||||
"""A memory pool that maps a request to its token locations."""
|
"""A memory pool that maps a request to its token locations."""
|
||||||
@@ -1232,6 +1259,10 @@ class HybridLinearKVPool(KVCache):
|
|||||||
)
|
)
|
||||||
return mamba_data_ptrs, mamba_data_lens, mamba_item_lens
|
return mamba_data_ptrs, mamba_data_lens, mamba_item_lens
|
||||||
|
|
||||||
|
def get_state_dim_per_tensor(self):
|
||||||
|
"""Get the sliceable dimension size for each mamba state tensor."""
|
||||||
|
return self.mamba_pool.get_state_dim_per_tensor()
|
||||||
|
|
||||||
def maybe_get_custom_mem_pool(self):
|
def maybe_get_custom_mem_pool(self):
|
||||||
return self.full_kv_pool.maybe_get_custom_mem_pool()
|
return self.full_kv_pool.maybe_get_custom_mem_pool()
|
||||||
|
|
||||||
|
|||||||
@@ -153,5 +153,79 @@ class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBa
|
|||||||
self.assertGreater(metrics["accuracy"], 0.93)
|
self.assertGreater(metrics["accuracy"], 0.93)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDisaggregationHybridAttentionMambaDPDecode(PDDisaggregationServerBase):
|
||||||
|
"""Test with prefill tp=2 and decode tp=2/dp=2 with dp-attention enabled."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
super().setUpClass()
|
||||||
|
cls.model = "Qwen/Qwen3-Next-80B-A3B-Instruct"
|
||||||
|
|
||||||
|
# Non blocking start servers
|
||||||
|
cls.start_prefill()
|
||||||
|
cls.start_decode()
|
||||||
|
|
||||||
|
# Block until both
|
||||||
|
cls.wait_server_ready(cls.prefill_url + "/health")
|
||||||
|
cls.wait_server_ready(cls.decode_url + "/health")
|
||||||
|
|
||||||
|
cls.launch_lb()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def start_prefill(cls):
|
||||||
|
prefill_args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--disaggregation-mode",
|
||||||
|
"prefill",
|
||||||
|
"--tp",
|
||||||
|
"2",
|
||||||
|
]
|
||||||
|
prefill_args += cls.transfer_backend + cls.rdma_devices
|
||||||
|
cls.process_prefill = popen_launch_pd_server(
|
||||||
|
cls.model,
|
||||||
|
cls.prefill_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=prefill_args,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def start_decode(cls):
|
||||||
|
decode_args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--disaggregation-mode",
|
||||||
|
"decode",
|
||||||
|
"--tp",
|
||||||
|
"2",
|
||||||
|
"--dp",
|
||||||
|
"2",
|
||||||
|
"--enable-dp-attention",
|
||||||
|
"--enable-dp-lm-head",
|
||||||
|
"--base-gpu-id",
|
||||||
|
"2",
|
||||||
|
]
|
||||||
|
decode_args += cls.transfer_backend + cls.rdma_devices
|
||||||
|
cls.process_decode = popen_launch_pd_server(
|
||||||
|
cls.model,
|
||||||
|
cls.decode_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=decode_args,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_gsm8k(self):
|
||||||
|
args = SimpleNamespace(
|
||||||
|
num_shots=5,
|
||||||
|
data_path=None,
|
||||||
|
num_questions=200,
|
||||||
|
max_new_tokens=512,
|
||||||
|
parallel=128,
|
||||||
|
host=f"http://{self.base_host}",
|
||||||
|
port=int(self.lb_port),
|
||||||
|
)
|
||||||
|
metrics = run_eval_few_shot_gsm8k(args)
|
||||||
|
print(f"Evaluation metrics: {metrics}")
|
||||||
|
|
||||||
|
self.assertGreater(metrics["accuracy"], 0.93)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user