[Bugfix] Fix Kimi-Linear state transfer across heterogeneous TP (#32262)
This commit is contained in:
@@ -224,6 +224,9 @@ class KimiLinearStateShape:
|
|||||||
# Conv tuples read (K-1, dim) — the overlapping dedup view would alias
|
# Conv tuples read (K-1, dim) — the overlapping dedup view would alias
|
||||||
# along the dim axis, so the dedup conv-intermediate layout must stay off.
|
# along the dim axis, so the dedup conv-intermediate layout must stay off.
|
||||||
disable_conv_window_dedup: bool = True
|
disable_conv_window_dedup: bool = True
|
||||||
|
# Per-slot conv tensors are [K-1, sharded_channels], unlike the usual
|
||||||
|
# [sharded_channels, K-1] layout.
|
||||||
|
conv_slice_axis: int = 1
|
||||||
|
|
||||||
num_heads: int
|
num_heads: int
|
||||||
head_dim: int
|
head_dim: int
|
||||||
@@ -231,6 +234,8 @@ class KimiLinearStateShape:
|
|||||||
head_k_dim: int
|
head_k_dim: int
|
||||||
conv_kernel: int
|
conv_kernel: int
|
||||||
num_spec: int
|
num_spec: int
|
||||||
|
# Full q/k/v dimensions. Each block is TP-sharded independently.
|
||||||
|
conv_shard_groups: Optional[List[int]] = None
|
||||||
# Number of key heads after TP sharding (== runtime ``H`` the KDA packed
|
# Number of key heads after TP sharding (== runtime ``H`` the KDA packed
|
||||||
# kernels infer from ``mixed_qkv``). Mirrors Mamba2StateShape; consumed by
|
# kernels infer from ``mixed_qkv``). Mirrors Mamba2StateShape; consumed by
|
||||||
# the ReplaySSM ring (k_cache) to size/stride exactly like the kernel.
|
# the ReplaySSM ring (k_cache) to size/stride exactly like the kernel.
|
||||||
@@ -278,6 +283,7 @@ class KimiLinearStateShape:
|
|||||||
head_k_dim=head_k_dim,
|
head_k_dim=head_k_dim,
|
||||||
conv_kernel=conv_kernel_size,
|
conv_kernel=conv_kernel_size,
|
||||||
num_spec=num_spec,
|
num_spec=num_spec,
|
||||||
|
conv_shard_groups=[proj_size, proj_k_size, proj_k_size],
|
||||||
num_k_heads_per_tp=num_k_heads_per_tp,
|
num_k_heads_per_tp=num_k_heads_per_tp,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -49,6 +49,8 @@ class KVArgs:
|
|||||||
state_item_lens: List[List[int]]
|
state_item_lens: List[List[int]]
|
||||||
# Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ.
|
# Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ.
|
||||||
state_dim_per_tensor: List[List[int]]
|
state_dim_per_tensor: List[List[int]]
|
||||||
|
# Number of rows before the slice axis in each per-slot state tensor.
|
||||||
|
state_slice_outer_counts: List[List[int]]
|
||||||
is_hybrid_mla_backend: bool
|
is_hybrid_mla_backend: bool
|
||||||
# Per-tensor conv sub-block dims (GDN: [key_dim, key_dim, value_dim]) so the
|
# Per-tensor conv sub-block dims (GDN: [key_dim, key_dim, value_dim]) so the
|
||||||
# scatter transfer can slice each independently head-sharded sub-block; None
|
# scatter transfer can slice each independently head-sharded sub-block; None
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ from sglang.srt.disaggregation.mooncake.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.utils import (
|
from sglang.srt.disaggregation.utils import (
|
||||||
DisaggregationMode,
|
DisaggregationMode,
|
||||||
compute_mamba_state_slice_blocks,
|
compute_mamba_state_slice_byte_blocks,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine
|
from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -948,7 +948,8 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
if sub_rank != 0:
|
if sub_rank != 0:
|
||||||
skip_kv = True
|
skip_kv = True
|
||||||
skip_state = True
|
# Hybrid-MLA KV is replicated across these source ranks, but
|
||||||
|
# TP-sharded state needs every rank for the aggregation path.
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.attn_cp_size > 1
|
self.attn_cp_size > 1
|
||||||
@@ -985,6 +986,12 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
src_conv_shard_groups = (
|
src_conv_shard_groups = (
|
||||||
src_conv_shard_groups[i] if i < len(src_conv_shard_groups) else []
|
src_conv_shard_groups[i] if i < len(src_conv_shard_groups) else []
|
||||||
)
|
)
|
||||||
|
src_slice_outer_counts = getattr(
|
||||||
|
self.kv_args, "state_slice_outer_counts", []
|
||||||
|
)
|
||||||
|
src_slice_outer_counts = (
|
||||||
|
src_slice_outer_counts[i] if i < len(src_slice_outer_counts) else []
|
||||||
|
)
|
||||||
if target_rank_registration_info is not None:
|
if target_rank_registration_info is not None:
|
||||||
dst_data_ptrs = (
|
dst_data_ptrs = (
|
||||||
target_rank_registration_info.dst_state_data_ptrs[i]
|
target_rank_registration_info.dst_state_data_ptrs[i]
|
||||||
@@ -1027,6 +1034,7 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
target_rank_registration_info.dst_tp_rank,
|
target_rank_registration_info.dst_tp_rank,
|
||||||
target_rank_registration_info.dst_attn_tp_size,
|
target_rank_registration_info.dst_attn_tp_size,
|
||||||
src_conv_shard_groups,
|
src_conv_shard_groups,
|
||||||
|
src_slice_outer_counts,
|
||||||
)
|
)
|
||||||
or rc
|
or rc
|
||||||
)
|
)
|
||||||
@@ -1166,6 +1174,7 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
dst_tp_rank: int,
|
dst_tp_rank: int,
|
||||||
dst_attn_tp_size: int,
|
dst_attn_tp_size: int,
|
||||||
src_state_conv_shard_groups: list = None,
|
src_state_conv_shard_groups: list = None,
|
||||||
|
src_state_slice_outer_counts: list[int] = None,
|
||||||
):
|
):
|
||||||
"""Transfer Mamba states with TP slice support.
|
"""Transfer Mamba states with TP slice support.
|
||||||
|
|
||||||
@@ -1207,45 +1216,41 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
src_dim = src_state_dim_per_tensor[i]
|
src_dim = src_state_dim_per_tensor[i]
|
||||||
dst_dim = dst_state_dim_per_tensor[i]
|
dst_dim = dst_state_dim_per_tensor[i]
|
||||||
|
|
||||||
# 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
|
|
||||||
|
|
||||||
conv_shard_groups = (
|
conv_shard_groups = (
|
||||||
src_state_conv_shard_groups[i]
|
src_state_conv_shard_groups[i]
|
||||||
if src_state_conv_shard_groups and i < len(src_state_conv_shard_groups)
|
if src_state_conv_shard_groups and i < len(src_state_conv_shard_groups)
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
# One block for single-axis states; three (q/k/v) for GDN conv_state
|
outer_count = (
|
||||||
# on the scatter path.
|
src_state_slice_outer_counts[i]
|
||||||
|
if src_state_slice_outer_counts
|
||||||
|
and i < len(src_state_slice_outer_counts)
|
||||||
|
else 1
|
||||||
|
)
|
||||||
for (
|
for (
|
||||||
src_dim_start,
|
src_offset,
|
||||||
dst_dim_start,
|
dst_offset,
|
||||||
num_dims_to_send,
|
bytes_to_send,
|
||||||
) in compute_mamba_state_slice_blocks(
|
) in compute_mamba_state_slice_byte_blocks(
|
||||||
|
src_item_len=src_item_len,
|
||||||
|
dst_item_len=dst_item_len,
|
||||||
src_dim=src_dim,
|
src_dim=src_dim,
|
||||||
dst_dim=dst_dim,
|
dst_dim=dst_dim,
|
||||||
|
outer_count=outer_count,
|
||||||
src_attn_tp_size=self.attn_tp_size,
|
src_attn_tp_size=self.attn_tp_size,
|
||||||
dst_attn_tp_size=dst_attn_tp_size,
|
dst_attn_tp_size=dst_attn_tp_size,
|
||||||
dst_tp_rank_in_group=dst_tp_rank_in_group,
|
dst_tp_rank_in_group=dst_tp_rank_in_group,
|
||||||
local_tp_rank_in_group=local_tp_rank_in_group,
|
local_tp_rank_in_group=local_tp_rank_in_group,
|
||||||
conv_shard_groups=conv_shard_groups,
|
conv_shard_groups=conv_shard_groups,
|
||||||
):
|
):
|
||||||
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
|
|
||||||
|
|
||||||
src_addr = (
|
src_addr = (
|
||||||
src_state_data_ptrs[i]
|
src_state_data_ptrs[i]
|
||||||
+ src_item_len * int(prefill_mamba_index[0])
|
+ src_item_len * int(prefill_mamba_index[0])
|
||||||
+ src_dim_offset
|
+ src_offset
|
||||||
)
|
)
|
||||||
dst_addr = (
|
dst_addr = (
|
||||||
dst_state_ptr
|
dst_state_ptr + dst_item_len * int(dst_mamba_index[0]) + dst_offset
|
||||||
+ dst_item_len * int(dst_mamba_index[0])
|
|
||||||
+ dst_dim_offset
|
|
||||||
)
|
)
|
||||||
|
|
||||||
transfer_blocks.append((src_addr, dst_addr, bytes_to_send))
|
transfer_blocks.append((src_addr, dst_addr, bytes_to_send))
|
||||||
|
|
||||||
return self._transfer_data(req.mooncake_session_id, transfer_blocks)
|
return self._transfer_data(req.mooncake_session_id, transfer_blocks)
|
||||||
|
|||||||
@@ -37,7 +37,7 @@ from sglang.srt.disaggregation.common.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.utils import (
|
from sglang.srt.disaggregation.utils import (
|
||||||
DisaggregationMode,
|
DisaggregationMode,
|
||||||
compute_mamba_state_slice_blocks,
|
compute_mamba_state_slice_byte_blocks,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
@@ -932,8 +932,11 @@ class NixlKVManager(CommonKVManager):
|
|||||||
f"({n_dst}) than prefill ({n_src}); unexpected geometry"
|
f"({n_dst}) than prefill ({n_src}); unexpected geometry"
|
||||||
)
|
)
|
||||||
decode_only_spec_dec = n_dst > n_src
|
decode_only_spec_dec = n_dst > n_src
|
||||||
|
if (
|
||||||
if self.is_mla_backend or peer_info.decode_tp_size == self.attn_tp_size:
|
self.is_mla_backend
|
||||||
|
or self.is_hybrid_mla_backend
|
||||||
|
or peer_info.decode_tp_size == self.attn_tp_size
|
||||||
|
):
|
||||||
dst_mem_kind = None
|
dst_mem_kind = None
|
||||||
try:
|
try:
|
||||||
dst_mem_kind = _homogeneous_kv_mem_kind(
|
dst_mem_kind = _homogeneous_kv_mem_kind(
|
||||||
@@ -1086,6 +1089,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
self.enable_staging
|
self.enable_staging
|
||||||
and staging_strategy is not None
|
and staging_strategy is not None
|
||||||
and not self.is_mla_backend
|
and not self.is_mla_backend
|
||||||
|
and not self.is_hybrid_mla_backend
|
||||||
and decode_tp_size != self.attn_tp_size
|
and decode_tp_size != self.attn_tp_size
|
||||||
and dst_info.staging is not None
|
and dst_info.staging is not None
|
||||||
)
|
)
|
||||||
@@ -1108,8 +1112,10 @@ class NixlKVManager(CommonKVManager):
|
|||||||
break
|
break
|
||||||
|
|
||||||
if kv_xfer_handle is None:
|
if kv_xfer_handle is None:
|
||||||
if self.is_mla_backend or (
|
if (
|
||||||
decode_tp_size == self.attn_tp_size
|
self.is_mla_backend
|
||||||
|
or self.is_hybrid_mla_backend
|
||||||
|
or decode_tp_size == self.attn_tp_size
|
||||||
):
|
):
|
||||||
if dst_info.kv_xfer_segments is None:
|
if dst_info.kv_xfer_segments is None:
|
||||||
if dst_info.dst_homogeneous_mem_kind is None:
|
if dst_info.dst_homogeneous_mem_kind is None:
|
||||||
@@ -1887,6 +1893,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
decode_tp_size: int,
|
decode_tp_size: int,
|
||||||
decode_tp_rank: int,
|
decode_tp_rank: int,
|
||||||
src_state_conv_shard_groups: list = None,
|
src_state_conv_shard_groups: list = None,
|
||||||
|
src_state_slice_outer_counts: list[int] = None,
|
||||||
):
|
):
|
||||||
"""Transfer Mamba states with TP slice support via RDMA.
|
"""Transfer Mamba states with TP slice support via RDMA.
|
||||||
|
|
||||||
@@ -1930,42 +1937,42 @@ class NixlKVManager(CommonKVManager):
|
|||||||
src_dim = src_state_dim_per_tensor[i]
|
src_dim = src_state_dim_per_tensor[i]
|
||||||
dst_dim = dst_state_dim_per_tensor[i]
|
dst_dim = dst_state_dim_per_tensor[i]
|
||||||
|
|
||||||
src_bytes_per_dim = src_item_len // src_dim
|
|
||||||
dst_bytes_per_dim = dst_item_len // dst_dim
|
|
||||||
|
|
||||||
conv_shard_groups = (
|
conv_shard_groups = (
|
||||||
src_state_conv_shard_groups[i]
|
src_state_conv_shard_groups[i]
|
||||||
if src_state_conv_shard_groups and i < len(src_state_conv_shard_groups)
|
if src_state_conv_shard_groups and i < len(src_state_conv_shard_groups)
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
# One block for single-axis states; three (q/k/v) for GDN conv_state
|
outer_count = (
|
||||||
# on the scatter path.
|
src_state_slice_outer_counts[i]
|
||||||
|
if src_state_slice_outer_counts
|
||||||
|
and i < len(src_state_slice_outer_counts)
|
||||||
|
else 1
|
||||||
|
)
|
||||||
for (
|
for (
|
||||||
src_dim_start,
|
src_offset,
|
||||||
dst_dim_start,
|
dst_offset,
|
||||||
num_dims_to_send,
|
bytes_to_send,
|
||||||
) in compute_mamba_state_slice_blocks(
|
) in compute_mamba_state_slice_byte_blocks(
|
||||||
|
src_item_len=src_item_len,
|
||||||
|
dst_item_len=dst_item_len,
|
||||||
src_dim=src_dim,
|
src_dim=src_dim,
|
||||||
dst_dim=dst_dim,
|
dst_dim=dst_dim,
|
||||||
|
outer_count=outer_count,
|
||||||
src_attn_tp_size=self.attn_tp_size,
|
src_attn_tp_size=self.attn_tp_size,
|
||||||
dst_attn_tp_size=decode_tp_size,
|
dst_attn_tp_size=decode_tp_size,
|
||||||
dst_tp_rank_in_group=dst_tp_rank_in_group,
|
dst_tp_rank_in_group=dst_tp_rank_in_group,
|
||||||
local_tp_rank_in_group=local_tp_rank_in_group,
|
local_tp_rank_in_group=local_tp_rank_in_group,
|
||||||
conv_shard_groups=conv_shard_groups,
|
conv_shard_groups=conv_shard_groups,
|
||||||
):
|
):
|
||||||
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
|
|
||||||
|
|
||||||
src_addr = (
|
src_addr = (
|
||||||
src_state_data_ptrs[i]
|
src_state_data_ptrs[i]
|
||||||
+ src_item_len * int(prefill_state_indices[0])
|
+ src_item_len * int(prefill_state_indices[0])
|
||||||
+ src_dim_offset
|
+ src_offset
|
||||||
)
|
)
|
||||||
dst_addr = (
|
dst_addr = (
|
||||||
dst_state_ptr
|
dst_state_ptr
|
||||||
+ dst_item_len * int(dst_state_indices[0])
|
+ dst_item_len * int(dst_state_indices[0])
|
||||||
+ dst_dim_offset
|
+ dst_offset
|
||||||
)
|
)
|
||||||
src_addrs.append((src_addr, bytes_to_send, self.kv_args.gpu_id))
|
src_addrs.append((src_addr, bytes_to_send, self.kv_args.gpu_id))
|
||||||
dst_addrs.append((dst_addr, bytes_to_send, dst_gpu_id))
|
dst_addrs.append((dst_addr, bytes_to_send, dst_gpu_id))
|
||||||
@@ -2010,6 +2017,9 @@ class NixlKVManager(CommonKVManager):
|
|||||||
src_state_conv_shard_groups = (
|
src_state_conv_shard_groups = (
|
||||||
getattr(self.kv_args, "state_conv_shard_groups", []) or []
|
getattr(self.kv_args, "state_conv_shard_groups", []) or []
|
||||||
)
|
)
|
||||||
|
src_state_slice_outer_counts = (
|
||||||
|
getattr(self.kv_args, "state_slice_outer_counts", []) or []
|
||||||
|
)
|
||||||
dst_state_item_lens = dst_state_item_lens or []
|
dst_state_item_lens = dst_state_item_lens or []
|
||||||
dst_state_dim_per_tensor = dst_state_dim_per_tensor or []
|
dst_state_dim_per_tensor = dst_state_dim_per_tensor or []
|
||||||
|
|
||||||
@@ -2030,6 +2040,11 @@ class NixlKVManager(CommonKVManager):
|
|||||||
if i < len(src_state_conv_shard_groups)
|
if i < len(src_state_conv_shard_groups)
|
||||||
else []
|
else []
|
||||||
)
|
)
|
||||||
|
src_outer_counts = (
|
||||||
|
src_state_slice_outer_counts[i]
|
||||||
|
if i < len(src_state_slice_outer_counts)
|
||||||
|
else []
|
||||||
|
)
|
||||||
dst_ptrs = dst_state_data_ptrs[i] if i < len(dst_state_data_ptrs) else []
|
dst_ptrs = dst_state_data_ptrs[i] if i < len(dst_state_data_ptrs) else []
|
||||||
dst_indices = dst_state_indices[i] if i < len(dst_state_indices) else []
|
dst_indices = dst_state_indices[i] if i < len(dst_state_indices) else []
|
||||||
dst_lens = dst_state_item_lens[i] if i < len(dst_state_item_lens) else []
|
dst_lens = dst_state_item_lens[i] if i < len(dst_state_item_lens) else []
|
||||||
@@ -2055,6 +2070,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
decode_tp_size,
|
decode_tp_size,
|
||||||
decode_tp_rank,
|
decode_tp_rank,
|
||||||
src_conv,
|
src_conv,
|
||||||
|
src_outer_counts,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
h = self._send_mamba_state(
|
h = self._send_mamba_state(
|
||||||
|
|||||||
@@ -809,6 +809,52 @@ def compute_mamba_state_slice_blocks(
|
|||||||
return blocks
|
return blocks
|
||||||
|
|
||||||
|
|
||||||
|
def compute_mamba_state_slice_byte_blocks(
|
||||||
|
*,
|
||||||
|
src_item_len: int,
|
||||||
|
dst_item_len: int,
|
||||||
|
src_dim: int,
|
||||||
|
dst_dim: int,
|
||||||
|
outer_count: int,
|
||||||
|
src_attn_tp_size: int,
|
||||||
|
dst_attn_tp_size: int,
|
||||||
|
dst_tp_rank_in_group: int,
|
||||||
|
local_tp_rank_in_group: int,
|
||||||
|
conv_shard_groups: Optional[List[int]] = None,
|
||||||
|
) -> List[Tuple[int, int, int]]:
|
||||||
|
"""Convert logical TP slices into physical byte blocks for one state slot.
|
||||||
|
|
||||||
|
``outer_count`` is one for the usual ``[slice_dim, ...]`` layout. Kimi
|
||||||
|
conv state is ``[K - 1, slice_dim]``, so each logical channel slice expands
|
||||||
|
into one byte block per convolution row.
|
||||||
|
"""
|
||||||
|
src_bytes_per_dim = src_item_len // (src_dim * outer_count)
|
||||||
|
dst_bytes_per_dim = dst_item_len // (dst_dim * outer_count)
|
||||||
|
logical_blocks = compute_mamba_state_slice_blocks(
|
||||||
|
src_dim=src_dim,
|
||||||
|
dst_dim=dst_dim,
|
||||||
|
src_attn_tp_size=src_attn_tp_size,
|
||||||
|
dst_attn_tp_size=dst_attn_tp_size,
|
||||||
|
dst_tp_rank_in_group=dst_tp_rank_in_group,
|
||||||
|
local_tp_rank_in_group=local_tp_rank_in_group,
|
||||||
|
conv_shard_groups=conv_shard_groups,
|
||||||
|
)
|
||||||
|
|
||||||
|
blocks = []
|
||||||
|
for outer_idx in range(outer_count):
|
||||||
|
src_row_offset = outer_idx * src_dim * src_bytes_per_dim
|
||||||
|
dst_row_offset = outer_idx * dst_dim * dst_bytes_per_dim
|
||||||
|
for src_dim_start, dst_dim_start, num_dims in logical_blocks:
|
||||||
|
blocks.append(
|
||||||
|
(
|
||||||
|
src_row_offset + src_dim_start * src_bytes_per_dim,
|
||||||
|
dst_row_offset + dst_dim_start * dst_bytes_per_dim,
|
||||||
|
num_dims * src_bytes_per_dim,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return blocks
|
||||||
|
|
||||||
|
|
||||||
def append_state_component(
|
def append_state_component(
|
||||||
kv_args: KVArgs,
|
kv_args: KVArgs,
|
||||||
state_type: StateType,
|
state_type: StateType,
|
||||||
@@ -817,6 +863,7 @@ def append_state_component(
|
|||||||
item_lens: List[int],
|
item_lens: List[int],
|
||||||
dim_per_tensor: Optional[List[int]] = None,
|
dim_per_tensor: Optional[List[int]] = None,
|
||||||
conv_shard_groups: Optional[List[Optional[List[int]]]] = None,
|
conv_shard_groups: Optional[List[Optional[List[int]]]] = None,
|
||||||
|
slice_outer_counts: Optional[List[int]] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Append one state component. Caller orders state_types consistently
|
"""Append one state component. Caller orders state_types consistently
|
||||||
on prefill and decode sides."""
|
on prefill and decode sides."""
|
||||||
@@ -826,6 +873,7 @@ def append_state_component(
|
|||||||
kv_args.state_item_lens.append(item_lens)
|
kv_args.state_item_lens.append(item_lens)
|
||||||
kv_args.state_dim_per_tensor.append(dim_per_tensor or [])
|
kv_args.state_dim_per_tensor.append(dim_per_tensor or [])
|
||||||
kv_args.state_conv_shard_groups.append(conv_shard_groups or [])
|
kv_args.state_conv_shard_groups.append(conv_shard_groups or [])
|
||||||
|
kv_args.state_slice_outer_counts.append(slice_outer_counts or [])
|
||||||
|
|
||||||
|
|
||||||
def setup_state_kv_args(
|
def setup_state_kv_args(
|
||||||
@@ -854,6 +902,7 @@ def setup_state_kv_args(
|
|||||||
kv_args.state_data_lens = []
|
kv_args.state_data_lens = []
|
||||||
kv_args.state_item_lens = []
|
kv_args.state_item_lens = []
|
||||||
kv_args.state_dim_per_tensor = []
|
kv_args.state_dim_per_tensor = []
|
||||||
|
kv_args.state_slice_outer_counts = []
|
||||||
kv_args.is_hybrid_mla_backend = False
|
kv_args.is_hybrid_mla_backend = False
|
||||||
kv_args.state_conv_shard_groups = []
|
kv_args.state_conv_shard_groups = []
|
||||||
|
|
||||||
@@ -917,6 +966,11 @@ def setup_state_kv_args(
|
|||||||
if hasattr(token_to_kv_pool, "get_state_conv_shard_groups")
|
if hasattr(token_to_kv_pool, "get_state_conv_shard_groups")
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
slice_outer_counts = (
|
||||||
|
token_to_kv_pool.get_state_slice_outer_counts()
|
||||||
|
if hasattr(token_to_kv_pool, "get_state_slice_outer_counts")
|
||||||
|
else None
|
||||||
|
)
|
||||||
append_state_component(
|
append_state_component(
|
||||||
kv_args,
|
kv_args,
|
||||||
StateType.MAMBA,
|
StateType.MAMBA,
|
||||||
@@ -925,6 +979,7 @@ def setup_state_kv_args(
|
|||||||
item_lens,
|
item_lens,
|
||||||
dim,
|
dim,
|
||||||
conv_shard_groups,
|
conv_shard_groups,
|
||||||
|
slice_outer_counts,
|
||||||
)
|
)
|
||||||
elif isinstance(token_to_kv_pool, (DSATokenToKVPool, NPUMLATokenToKVPool)):
|
elif isinstance(token_to_kv_pool, (DSATokenToKVPool, NPUMLATokenToKVPool)):
|
||||||
if draft_token_to_kv_pool is not None and isinstance(
|
if draft_token_to_kv_pool is not None and isinstance(
|
||||||
@@ -1041,6 +1096,11 @@ def setup_state_kv_args(
|
|||||||
if hasattr(req_to_token_pool, "get_state_conv_shard_groups")
|
if hasattr(req_to_token_pool, "get_state_conv_shard_groups")
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
slice_outer_counts = (
|
||||||
|
req_to_token_pool.get_state_slice_outer_counts()
|
||||||
|
if hasattr(req_to_token_pool, "get_state_slice_outer_counts")
|
||||||
|
else None
|
||||||
|
)
|
||||||
append_state_component(
|
append_state_component(
|
||||||
kv_args,
|
kv_args,
|
||||||
StateType.MAMBA,
|
StateType.MAMBA,
|
||||||
@@ -1049,6 +1109,7 @@ def setup_state_kv_args(
|
|||||||
item_lens,
|
item_lens,
|
||||||
dim,
|
dim,
|
||||||
conv_shard_groups,
|
conv_shard_groups,
|
||||||
|
slice_outer_counts,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -811,6 +811,7 @@ class MambaPool:
|
|||||||
# Full (unsharded) conv sub-block dims for PD transfer across different
|
# Full (unsharded) conv sub-block dims for PD transfer across different
|
||||||
# attn_tp_size (GDN: [key_dim, key_dim, value_dim]); None otherwise.
|
# attn_tp_size (GDN: [key_dim, key_dim, value_dim]); None otherwise.
|
||||||
self.conv_shard_groups = getattr(cache_params.shape, "conv_shard_groups", None)
|
self.conv_shard_groups = getattr(cache_params.shape, "conv_shard_groups", None)
|
||||||
|
self.conv_slice_axis = getattr(cache_params.shape, "conv_slice_axis", 0)
|
||||||
|
|
||||||
def get_speculative_mamba2_params_all_layers(self) -> SpeculativeState:
|
def get_speculative_mamba2_params_all_layers(self) -> SpeculativeState:
|
||||||
assert isinstance(self.mamba_cache, self.SpeculativeState)
|
assert isinstance(self.mamba_cache, self.SpeculativeState)
|
||||||
@@ -984,39 +985,34 @@ class MambaPool:
|
|||||||
)
|
)
|
||||||
current_platform.synchronize()
|
current_platform.synchronize()
|
||||||
|
|
||||||
|
_NON_TRANSFER_STATE_FIELDS = frozenset(
|
||||||
|
{
|
||||||
|
"intermediate_ssm",
|
||||||
|
"intermediate_conv_window",
|
||||||
|
"replayssm_d",
|
||||||
|
"replayssm_k",
|
||||||
|
"replayssm_g",
|
||||||
|
"replayssm_rawv",
|
||||||
|
"replayssm_rawk",
|
||||||
|
"replayssm_beta",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
def _iter_transfer_state_tensors(self):
|
||||||
|
"""Yield transferable state tensors with their per-slot slice axis."""
|
||||||
|
for field, value in vars(self.mamba_cache).items():
|
||||||
|
if field in self._NON_TRANSFER_STATE_FIELDS or value is None:
|
||||||
|
continue
|
||||||
|
tensors = value if isinstance(value, list) else [value]
|
||||||
|
slice_axis = self.conv_slice_axis if field == "conv" else 0
|
||||||
|
for state_tensor in tensors:
|
||||||
|
yield field, state_tensor, slice_axis
|
||||||
|
|
||||||
def get_contiguous_buf_infos(self):
|
def get_contiguous_buf_infos(self):
|
||||||
"""
|
"""Get transferable state buffer information for RDMA registration."""
|
||||||
Get buffer info for RDMA registration.
|
|
||||||
Only returns conv and temporal state buffers, excluding intermediate buffers
|
|
||||||
used for speculative decoding (intermediate_ssm, intermediate_conv_window).
|
|
||||||
"""
|
|
||||||
state_tensors = []
|
|
||||||
for field in vars(self.mamba_cache):
|
|
||||||
# Skip intermediate buffers used only for speculative decoding
|
|
||||||
# These buffers have different size (spec_state_size + 1) and should not be transferred
|
|
||||||
if field in ("intermediate_ssm", "intermediate_conv_window"):
|
|
||||||
continue
|
|
||||||
# Skip GDN ReplaySSM ring buffers: they are derived/transient decode
|
|
||||||
# scratch, not part of the persistent transferable state.
|
|
||||||
if field in (
|
|
||||||
"replayssm_d",
|
|
||||||
"replayssm_k",
|
|
||||||
"replayssm_g",
|
|
||||||
"replayssm_rawv",
|
|
||||||
"replayssm_rawk",
|
|
||||||
"replayssm_beta",
|
|
||||||
):
|
|
||||||
continue
|
|
||||||
value = getattr(self.mamba_cache, field)
|
|
||||||
if value is None:
|
|
||||||
continue
|
|
||||||
if isinstance(value, list):
|
|
||||||
state_tensors.extend(value)
|
|
||||||
else:
|
|
||||||
state_tensors.append(value)
|
|
||||||
data_ptrs, data_lens, item_lens = [], [], []
|
data_ptrs, data_lens, item_lens = [], [], []
|
||||||
|
|
||||||
for _, state_tensor in enumerate(state_tensors):
|
for _, state_tensor, _ in self._iter_transfer_state_tensors():
|
||||||
data_ptrs += [
|
data_ptrs += [
|
||||||
state_tensor[i].data_ptr() for i in range(self.num_mamba_layers)
|
state_tensor[i].data_ptr() for i in range(self.num_mamba_layers)
|
||||||
]
|
]
|
||||||
@@ -1029,45 +1025,27 @@ class MambaPool:
|
|||||||
def get_state_dim_per_tensor(self):
|
def get_state_dim_per_tensor(self):
|
||||||
"""Get the sliceable dimension size for each state tensor.
|
"""Get the sliceable dimension size for each state tensor.
|
||||||
|
|
||||||
For mamba state, the layout is:
|
The slice axis is tensor-specific: normally the first per-slot axis,
|
||||||
- conv_state: [num_layers, size+1, conv_dim/tp, conv_kernel-1]
|
while Kimi conv state uses the second per-slot axis.
|
||||||
- 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):
|
|
||||||
# Mirror the exclusions in get_contiguous_buf_infos so the returned
|
|
||||||
# dims line up element-wise with the RDMA buffer list.
|
|
||||||
if field in (
|
|
||||||
"intermediate_ssm",
|
|
||||||
"intermediate_conv_window",
|
|
||||||
"replayssm_d",
|
|
||||||
"replayssm_k",
|
|
||||||
"replayssm_g",
|
|
||||||
"replayssm_rawv",
|
|
||||||
"replayssm_rawk",
|
|
||||||
"replayssm_beta",
|
|
||||||
):
|
|
||||||
continue
|
|
||||||
value = getattr(self.mamba_cache, field)
|
|
||||||
if value is None:
|
|
||||||
continue
|
|
||||||
if isinstance(value, list):
|
|
||||||
state_tensors.extend(value)
|
|
||||||
else:
|
|
||||||
state_tensors.append(value)
|
|
||||||
|
|
||||||
dim_per_tensor = []
|
dim_per_tensor = []
|
||||||
for state_tensor in state_tensors:
|
for _, state_tensor, slice_axis in self._iter_transfer_state_tensors():
|
||||||
# state_tensor shape: [num_layers, size+1, sliceable_dim, ...]
|
# state_tensor shape: [num_layers, size+1, sliceable_dim, ...]
|
||||||
# The sliceable dimension is at index 2 (after num_layers and size)
|
# Kimi conv state transposes the two per-slot axes to [K-1, dim].
|
||||||
sliceable_dim = state_tensor.shape[2]
|
axis = 2 + slice_axis
|
||||||
|
sliceable_dim = state_tensor.shape[axis]
|
||||||
# Repeat for each layer since we have per-layer data_ptrs
|
# Repeat for each layer since we have per-layer data_ptrs
|
||||||
dim_per_tensor += [sliceable_dim] * self.num_mamba_layers
|
dim_per_tensor += [sliceable_dim] * self.num_mamba_layers
|
||||||
return dim_per_tensor
|
return dim_per_tensor
|
||||||
|
|
||||||
|
def get_state_slice_outer_counts(self):
|
||||||
|
"""Get the number of rows preceding each tensor's TP slice axis."""
|
||||||
|
outer_counts = []
|
||||||
|
for _, state_tensor, slice_axis in self._iter_transfer_state_tensors():
|
||||||
|
outer_count = math.prod(state_tensor.shape[2 : 2 + slice_axis])
|
||||||
|
outer_counts += [outer_count] * self.num_mamba_layers
|
||||||
|
return outer_counts
|
||||||
|
|
||||||
def get_state_conv_shard_groups(self):
|
def get_state_conv_shard_groups(self):
|
||||||
"""Per-tensor conv sub-block dims, aligned element-wise with
|
"""Per-tensor conv sub-block dims, aligned element-wise with
|
||||||
get_state_dim_per_tensor().
|
get_state_dim_per_tensor().
|
||||||
@@ -1080,29 +1058,14 @@ class MambaPool:
|
|||||||
those tensors keep the single contiguous slice.
|
those tensors keep the single contiguous slice.
|
||||||
"""
|
"""
|
||||||
subdims_per_tensor = []
|
subdims_per_tensor = []
|
||||||
for field in vars(self.mamba_cache):
|
for field, _, _ in self._iter_transfer_state_tensors():
|
||||||
# Mirror the exclusions in get_state_dim_per_tensor so the returned
|
# Only conv_state carries a q/k/v decomposition.
|
||||||
# sub-dims line up element-wise with the RDMA buffer list.
|
subdims = (
|
||||||
if field in (
|
list(self.conv_shard_groups)
|
||||||
"intermediate_ssm",
|
if field == "conv" and self.conv_shard_groups is not None
|
||||||
"intermediate_conv_window",
|
else None
|
||||||
"replayssm_d",
|
)
|
||||||
"replayssm_k",
|
subdims_per_tensor += [subdims] * self.num_mamba_layers
|
||||||
"replayssm_g",
|
|
||||||
):
|
|
||||||
continue
|
|
||||||
value = getattr(self.mamba_cache, field)
|
|
||||||
if value is None:
|
|
||||||
continue
|
|
||||||
tensors = value if isinstance(value, list) else [value]
|
|
||||||
for _ in tensors:
|
|
||||||
# Only conv_state carries a q/k/v decomposition.
|
|
||||||
subdims = (
|
|
||||||
list(self.conv_shard_groups)
|
|
||||||
if field == "conv" and self.conv_shard_groups is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
subdims_per_tensor += [subdims] * self.num_mamba_layers
|
|
||||||
return subdims_per_tensor
|
return subdims_per_tensor
|
||||||
|
|
||||||
|
|
||||||
@@ -1340,6 +1303,9 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
def get_state_dim_per_tensor(self):
|
def get_state_dim_per_tensor(self):
|
||||||
return self.mamba_pool.get_state_dim_per_tensor()
|
return self.mamba_pool.get_state_dim_per_tensor()
|
||||||
|
|
||||||
|
def get_state_slice_outer_counts(self):
|
||||||
|
return self.mamba_pool.get_state_slice_outer_counts()
|
||||||
|
|
||||||
def get_state_conv_shard_groups(self):
|
def get_state_conv_shard_groups(self):
|
||||||
return self.mamba_pool.get_state_conv_shard_groups()
|
return self.mamba_pool.get_state_conv_shard_groups()
|
||||||
|
|
||||||
@@ -3637,6 +3603,10 @@ class HybridLinearKVPool(KVCache):
|
|||||||
"""Get the sliceable dimension size for each mamba state tensor."""
|
"""Get the sliceable dimension size for each mamba state tensor."""
|
||||||
return self.mamba_pool.get_state_dim_per_tensor()
|
return self.mamba_pool.get_state_dim_per_tensor()
|
||||||
|
|
||||||
|
def get_state_slice_outer_counts(self):
|
||||||
|
"""Get the row count preceding each mamba state slice axis."""
|
||||||
|
return self.mamba_pool.get_state_slice_outer_counts()
|
||||||
|
|
||||||
def get_state_conv_shard_groups(self):
|
def get_state_conv_shard_groups(self):
|
||||||
"""Per-tensor conv sub-block dims (GDN) aligned with the state list."""
|
"""Per-tensor conv sub-block dims (GDN) aligned with the state list."""
|
||||||
return self.mamba_pool.get_state_conv_shard_groups()
|
return self.mamba_pool.get_state_conv_shard_groups()
|
||||||
|
|||||||
@@ -147,6 +147,7 @@ class MambaSubPoolSpec(SubPoolSpec):
|
|||||||
conv_dtype: torch.dtype
|
conv_dtype: torch.dtype
|
||||||
temporal_state_shape: Tuple[int, ...]
|
temporal_state_shape: Tuple[int, ...]
|
||||||
temporal_dtype: torch.dtype
|
temporal_dtype: torch.dtype
|
||||||
|
conv_slice_axis: int = 0
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
super().__post_init__()
|
super().__post_init__()
|
||||||
@@ -538,6 +539,12 @@ class UnifiedMambaPool(MambaPool):
|
|||||||
self.linear_replayssm_cache_len = 16
|
self.linear_replayssm_cache_len = 16
|
||||||
self.replayssm_write_pos = None
|
self.replayssm_write_pos = None
|
||||||
self.replayssm_is_kda = False
|
self.replayssm_is_kda = False
|
||||||
|
self.enable_gdn_replayssm_spec = False
|
||||||
|
self.replayssm_cache_base = None
|
||||||
|
self.replayssm_is_flush = None
|
||||||
|
self.debug_memory_pool = False
|
||||||
|
self.conv_shard_groups = None
|
||||||
|
self.conv_slice_axis = spec.conv_slice_axis
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
conv_views[0].shape[0] == self.num_mamba_layers
|
conv_views[0].shape[0] == self.num_mamba_layers
|
||||||
@@ -874,6 +881,7 @@ def init_unified_mamba_pools(
|
|||||||
conv_dtype=cp.dtype.conv,
|
conv_dtype=cp.dtype.conv,
|
||||||
temporal_state_shape=tuple(int(x) for x in cp.shape.temporal),
|
temporal_state_shape=tuple(int(x) for x in cp.shape.temporal),
|
||||||
temporal_dtype=cp.dtype.temporal,
|
temporal_dtype=cp.dtype.temporal,
|
||||||
|
conv_slice_axis=getattr(cp.shape, "conv_slice_axis", 0),
|
||||||
grow_direction="up",
|
grow_direction="up",
|
||||||
)
|
)
|
||||||
total_bytes = (
|
total_bytes = (
|
||||||
|
|||||||
@@ -185,7 +185,7 @@ class KimiDeltaAttention(nn.Module):
|
|||||||
self.num_k_heads = config.linear_attn_config["num_heads"]
|
self.num_k_heads = config.linear_attn_config["num_heads"]
|
||||||
self.num_v_heads = config.linear_attn_config["num_heads"]
|
self.num_v_heads = config.linear_attn_config["num_heads"]
|
||||||
self.head_k_dim = config.linear_attn_config["head_dim"]
|
self.head_k_dim = config.linear_attn_config["head_dim"]
|
||||||
self.head_v_dim = config.v_head_dim
|
self.head_v_dim = config.linear_attn_config["head_dim"]
|
||||||
self.layer_idx = layer_idx
|
self.layer_idx = layer_idx
|
||||||
self.prefix = prefix
|
self.prefix = prefix
|
||||||
assert self.num_heads % self.tp_size == 0
|
assert self.num_heads % self.tp_size == 0
|
||||||
|
|||||||
@@ -50,6 +50,9 @@ class PDDisaggregationServerBase(CustomTestCase):
|
|||||||
capture_per_side_logs: ClassVar[bool] = False
|
capture_per_side_logs: ClassVar[bool] = False
|
||||||
extra_prefill_env: ClassVar[dict[str, str]] = {}
|
extra_prefill_env: ClassVar[dict[str, str]] = {}
|
||||||
extra_decode_env: ClassVar[dict[str, str]] = {}
|
extra_decode_env: ClassVar[dict[str, str]] = {}
|
||||||
|
prefill_tp_size: ClassVar[int] = 1
|
||||||
|
decode_tp_size: ClassVar[int] = 1
|
||||||
|
decode_base_gpu_id: ClassVar[int] = 1
|
||||||
_prefill_stdout_buf: ClassVar[Optional[io.StringIO]] = None
|
_prefill_stdout_buf: ClassVar[Optional[io.StringIO]] = None
|
||||||
_prefill_stderr_buf: ClassVar[Optional[io.StringIO]] = None
|
_prefill_stderr_buf: ClassVar[Optional[io.StringIO]] = None
|
||||||
_decode_stdout_buf: ClassVar[Optional[io.StringIO]] = None
|
_decode_stdout_buf: ClassVar[Optional[io.StringIO]] = None
|
||||||
@@ -114,7 +117,7 @@ class PDDisaggregationServerBase(CustomTestCase):
|
|||||||
"--disaggregation-bootstrap-port",
|
"--disaggregation-bootstrap-port",
|
||||||
cls.bootstrap_port,
|
cls.bootstrap_port,
|
||||||
"--tp",
|
"--tp",
|
||||||
"1",
|
str(cls.prefill_tp_size),
|
||||||
] + list(cls.extra_prefill_args)
|
] + list(cls.extra_prefill_args)
|
||||||
prefill_args += cls.transfer_backend + cls.rdma_devices
|
prefill_args += cls.transfer_backend + cls.rdma_devices
|
||||||
cls.process_prefill = popen_launch_pd_server(
|
cls.process_prefill = popen_launch_pd_server(
|
||||||
@@ -139,9 +142,9 @@ class PDDisaggregationServerBase(CustomTestCase):
|
|||||||
"--disaggregation-bootstrap-port",
|
"--disaggregation-bootstrap-port",
|
||||||
cls.bootstrap_port,
|
cls.bootstrap_port,
|
||||||
"--tp",
|
"--tp",
|
||||||
"1",
|
str(cls.decode_tp_size),
|
||||||
"--base-gpu-id",
|
"--base-gpu-id",
|
||||||
"1",
|
str(cls.decode_base_gpu_id),
|
||||||
] + list(cls.extra_decode_args)
|
] + list(cls.extra_decode_args)
|
||||||
decode_args += cls.transfer_backend + cls.rdma_devices
|
decode_args += cls.transfer_backend + cls.rdma_devices
|
||||||
cls.process_decode = popen_launch_pd_server(
|
cls.process_decode = popen_launch_pd_server(
|
||||||
|
|||||||
@@ -0,0 +1,105 @@
|
|||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||||
|
PDDisaggregationServerBase,
|
||||||
|
assert_process_healthy,
|
||||||
|
)
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=240, stage="base-c", runner_config="4-gpu-h100")
|
||||||
|
|
||||||
|
KIMI_LINEAR_MODEL = "yujiepan/kimi-linear-tiny-random"
|
||||||
|
SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"}
|
||||||
|
SERVER_ARGS = [
|
||||||
|
"--skip-tokenizer-init",
|
||||||
|
"--random-seed",
|
||||||
|
"1",
|
||||||
|
"--enable-deterministic-inference",
|
||||||
|
"--max-mamba-cache-size",
|
||||||
|
"32",
|
||||||
|
"--max-total-tokens",
|
||||||
|
"4096",
|
||||||
|
"--cuda-graph-backend-decode",
|
||||||
|
"disabled",
|
||||||
|
"--cuda-graph-backend-prefill",
|
||||||
|
"disabled",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase):
|
||||||
|
prefill_tp_size = 2
|
||||||
|
decode_tp_size = 1
|
||||||
|
decode_base_gpu_id = 2
|
||||||
|
extra_prefill_args = SERVER_ARGS
|
||||||
|
extra_decode_args = SERVER_ARGS
|
||||||
|
extra_prefill_env = SERVER_ENV
|
||||||
|
extra_decode_env = SERVER_ENV
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
super().setUpClass()
|
||||||
|
cls.model = KIMI_LINEAR_MODEL
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def generate(base_url):
|
||||||
|
response = requests.post(
|
||||||
|
base_url + "/generate",
|
||||||
|
json={
|
||||||
|
"input_ids": [1] + [100 + i % 1000 for i in range(256)],
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0,
|
||||||
|
"max_new_tokens": 4,
|
||||||
|
"ignore_eos": True,
|
||||||
|
},
|
||||||
|
"return_logprob": True,
|
||||||
|
"top_logprobs_num": 5,
|
||||||
|
},
|
||||||
|
timeout=120,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return response.json()["meta_info"]
|
||||||
|
|
||||||
|
def test_logprob_parity(self):
|
||||||
|
baseline = popen_launch_server(
|
||||||
|
self.model,
|
||||||
|
self.lb_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=["--tp-size", "2", "--trust-remote-code"] + SERVER_ARGS,
|
||||||
|
env=SERVER_ENV,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
reference = self.generate(self.lb_url)
|
||||||
|
finally:
|
||||||
|
kill_process_tree(baseline.pid, wait_timeout=60)
|
||||||
|
time.sleep(5)
|
||||||
|
|
||||||
|
self.launch_all()
|
||||||
|
disaggregated = self.generate(self.lb_url)
|
||||||
|
|
||||||
|
reference_logprobs = reference["output_token_logprobs"]
|
||||||
|
disaggregated_logprobs = disaggregated["output_token_logprobs"]
|
||||||
|
self.assertEqual(
|
||||||
|
[item[1] for item in reference_logprobs],
|
||||||
|
[item[1] for item in disaggregated_logprobs],
|
||||||
|
)
|
||||||
|
self.assertEqual(len(reference_logprobs), 4)
|
||||||
|
for reference_item, disaggregated_item in zip(
|
||||||
|
reference_logprobs, disaggregated_logprobs
|
||||||
|
):
|
||||||
|
self.assertAlmostEqual(reference_item[0], disaggregated_item[0], delta=0.05)
|
||||||
|
|
||||||
|
assert_process_healthy(self, "load balancer", self.process_lb, self.lb_url)
|
||||||
|
assert_process_healthy(self, "prefill", self.process_prefill, self.prefill_url)
|
||||||
|
assert_process_healthy(self, "decode", self.process_decode, self.decode_url)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user