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:
@@ -24,6 +24,8 @@ class KVArgs:
|
||||
state_data_lens: List[int]
|
||||
state_item_lens: List[int]
|
||||
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_traffic_class: str
|
||||
gpu_id: int
|
||||
|
||||
@@ -292,6 +292,11 @@ class DecodePreallocQueue:
|
||||
kv_args.state_type = "swa"
|
||||
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
|
||||
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):
|
||||
kv_args.state_type = "nsa"
|
||||
else:
|
||||
|
||||
@@ -114,6 +114,9 @@ class KVArgsRegisterInfo:
|
||||
dst_tp_rank: int
|
||||
dst_attn_tp_size: 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
|
||||
def from_zmq(cls, msg: List[bytes]):
|
||||
@@ -128,6 +131,16 @@ class KVArgsRegisterInfo:
|
||||
dst_tp_rank=int(msg[7].decode("ascii")),
|
||||
dst_attn_tp_size=int(msg[8].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],
|
||||
dst_state_data_ptrs: list[int],
|
||||
executor: concurrent.futures.ThreadPoolExecutor,
|
||||
target_rank_registration_info: Optional[KVArgsRegisterInfo] = None,
|
||||
):
|
||||
"""Send state or extra pool data with type-specific handling."""
|
||||
state_type = getattr(self.kv_args, "state_type", "none")
|
||||
|
||||
if state_type == "mamba":
|
||||
return self._send_mamba_state(
|
||||
req,
|
||||
prefill_state_indices,
|
||||
dst_state_data_ptrs,
|
||||
)
|
||||
# Check if we need slice transfer for different TP sizes
|
||||
if (
|
||||
target_rank_registration_info is not None
|
||||
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"]:
|
||||
# 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
|
||||
prefill_state_indices = np.array(prefill_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)
|
||||
|
||||
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(
|
||||
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.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(
|
||||
req,
|
||||
kv_chunk.state_indices,
|
||||
target_rank_registration_info.dst_state_data_ptrs,
|
||||
executor,
|
||||
target_rank_registration_info,
|
||||
)
|
||||
|
||||
# Only the last chunk we need to send the aux data
|
||||
@@ -1204,6 +1319,17 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
||||
packed_state_data_ptrs = b"".join(
|
||||
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
|
||||
tp_rank = self.kv_mgr.kv_args.engine_rank
|
||||
kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0]
|
||||
@@ -1225,6 +1351,8 @@ class MooncakeKVReceiver(CommonKVReceiver):
|
||||
dst_tp_rank,
|
||||
dst_attn_tp_size,
|
||||
dst_kv_item_len,
|
||||
packed_state_item_lens,
|
||||
packed_state_dim_per_tensor,
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@@ -161,6 +161,11 @@ class PrefillBootstrapQueue:
|
||||
kv_args.state_type = "swa"
|
||||
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
|
||||
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):
|
||||
kv_args.state_type = "nsa"
|
||||
else:
|
||||
|
||||
@@ -374,6 +374,33 @@ class MambaPool:
|
||||
]
|
||||
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):
|
||||
"""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
|
||||
|
||||
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):
|
||||
return self.full_kv_pool.maybe_get_custom_mem_pool()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user