[Qwen3.8-Next] Add PD state transfer for Flash Next (#36651)
This commit is contained in:
@@ -33,15 +33,23 @@ from sglang.srt.disaggregation.mooncake.conn import (
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import (
|
||||
MetadataBuffers,
|
||||
build_transfer_entry_pairs,
|
||||
compute_mamba_state_slice_byte_blocks,
|
||||
get_dsv4_c4_state_indices,
|
||||
get_dsv4_c128_state_indices,
|
||||
get_qsa_pending_state_indices,
|
||||
setup_state_kv_args,
|
||||
should_send_replicated_state,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.dsa.utils import should_use_dsa_fused_topk
|
||||
from sglang.srt.managers.overlap_utils import FutureMap, RelayPayload
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.mem_cache.qsa_kv_pool import (
|
||||
QSA_ROPE_STATE_LAYER_ID,
|
||||
QSATokenToKVPool,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.srt.speculative.eagle_disaggregation import (
|
||||
build_eagle_disagg_draft_input,
|
||||
@@ -189,6 +197,132 @@ class TestCPReplicatedStateTransfer(unittest.TestCase):
|
||||
)
|
||||
|
||||
|
||||
class TestQwen4StateWire(unittest.TestCase):
|
||||
def test_qsa_pending_payload_uses_nested_request_pool_row(self):
|
||||
req = SimpleNamespace(kv=ReqKvInfo(req_pool_idx=7))
|
||||
|
||||
np.testing.assert_array_equal(
|
||||
get_qsa_pending_state_indices(req),
|
||||
np.array([7], dtype=np.int32),
|
||||
)
|
||||
|
||||
def test_qsa_registers_request_ring_and_page_state_separately(self):
|
||||
pool = object.__new__(QSATokenToKVPool)
|
||||
pool.full_kv_pool = object()
|
||||
pool.get_state_buf_infos = lambda: ([10], [100], [20])
|
||||
pool.get_state_dim_per_tensor = lambda: [4]
|
||||
pool.get_state_conv_shard_groups = lambda: [None]
|
||||
pool.get_state_slice_outer_counts = lambda: [1]
|
||||
pool.get_state_layer_ids = lambda: [2]
|
||||
pool.page_size = 4
|
||||
pool.qsa_compress_ratio = 2
|
||||
pool.qsa_compressed_page_size = 2
|
||||
pool.full_attention_layer_id_mapping = {24: 0}
|
||||
pool.qsa_key_state_buffer_pool = [torch.zeros((6, 1, 8), dtype=torch.bfloat16)]
|
||||
pool.qsa_rope_position_buffer = torch.zeros((6, 3), dtype=torch.int64)
|
||||
pool.qsa_compressed_k_buffer_pool = [
|
||||
torch.zeros((6, 1, 8), dtype=torch.bfloat16)
|
||||
]
|
||||
|
||||
kv_args = SimpleNamespace()
|
||||
setup_state_kv_args(kv_args, pool)
|
||||
|
||||
self.assertEqual(
|
||||
kv_args.state_types,
|
||||
[StateType.MAMBA, StateType.QSA_PENDING, StateType.QSA_COMPRESSED],
|
||||
)
|
||||
# Pending entries are whole two-row request rings; compressed-K remains
|
||||
# a two-row compressed page corresponding to one four-token KV page.
|
||||
self.assertEqual(kv_args.state_item_lens[1:], [[32, 48], [32]])
|
||||
self.assertEqual(
|
||||
kv_args.state_layer_ids[1:],
|
||||
[[24, QSA_ROPE_STATE_LAYER_ID], [24]],
|
||||
)
|
||||
|
||||
def test_qsa_stage_without_qsa_layers_does_not_register_rope_ring(self):
|
||||
pool = object.__new__(QSATokenToKVPool)
|
||||
pool.full_kv_pool = object()
|
||||
pool.get_state_buf_infos = lambda: ([10], [100], [20])
|
||||
pool.get_state_dim_per_tensor = lambda: [4]
|
||||
pool.get_state_conv_shard_groups = lambda: [None]
|
||||
pool.get_state_slice_outer_counts = lambda: [1]
|
||||
pool.get_state_layer_ids = lambda: [2]
|
||||
pool.page_size = 4
|
||||
pool.qsa_compress_ratio = 2
|
||||
pool.qsa_compressed_page_size = 2
|
||||
pool.full_attention_layer_id_mapping = {}
|
||||
pool.qsa_key_state_buffer_pool = []
|
||||
pool.qsa_rope_position_buffer = torch.zeros((6, 3), dtype=torch.int64)
|
||||
pool.qsa_compressed_k_buffer_pool = []
|
||||
|
||||
kv_args = SimpleNamespace()
|
||||
setup_state_kv_args(kv_args, pool)
|
||||
|
||||
# Keep the component slots aligned across PP stages, but expose no QSA
|
||||
# buffers or layer ids from a stage that cannot produce their contents.
|
||||
self.assertEqual(
|
||||
kv_args.state_types,
|
||||
[StateType.MAMBA, StateType.QSA_PENDING, StateType.QSA_COMPRESSED],
|
||||
)
|
||||
self.assertEqual(kv_args.state_data_ptrs[1:], [[], []])
|
||||
self.assertEqual(kv_args.state_data_lens[1:], [[], []])
|
||||
self.assertEqual(kv_args.state_item_lens[1:], [[], []])
|
||||
self.assertEqual(kv_args.state_layer_ids[1:], [[], []])
|
||||
|
||||
def test_compact_qsa_entries_map_by_global_layer_id(self):
|
||||
self.assertEqual(
|
||||
build_transfer_entry_pairs(
|
||||
[24, QSA_ROPE_STATE_LAYER_ID],
|
||||
[0, 12, 24, QSA_ROPE_STATE_LAYER_ID],
|
||||
2,
|
||||
4,
|
||||
),
|
||||
[(0, 2), (1, 3)],
|
||||
)
|
||||
|
||||
def test_replicated_state_tp_policy(self):
|
||||
for src_tp, dst_tp, rank, expected in (
|
||||
(4, 1, 0, True),
|
||||
(4, 1, 1, False),
|
||||
(1, 4, 0, True),
|
||||
(4, 4, 3, True),
|
||||
):
|
||||
with self.subTest(src_tp=src_tp, dst_tp=dst_tp, rank=rank):
|
||||
self.assertEqual(
|
||||
should_send_replicated_state(
|
||||
src_attn_tp_size=src_tp,
|
||||
dst_attn_tp_size=dst_tp,
|
||||
local_tp_rank_in_group=rank,
|
||||
),
|
||||
expected,
|
||||
)
|
||||
|
||||
common = dict(
|
||||
src_item_len=96,
|
||||
dst_item_len=96,
|
||||
src_dim=0,
|
||||
dst_dim=0,
|
||||
outer_count=1,
|
||||
src_attn_tp_size=4,
|
||||
dst_attn_tp_size=1,
|
||||
dst_tp_rank_in_group=0,
|
||||
)
|
||||
self.assertEqual(
|
||||
compute_mamba_state_slice_byte_blocks(**common, local_tp_rank_in_group=0),
|
||||
[(0, 0, 96)],
|
||||
)
|
||||
self.assertEqual(
|
||||
compute_mamba_state_slice_byte_blocks(**common, local_tp_rank_in_group=1),
|
||||
[],
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "must divide"):
|
||||
should_send_replicated_state(
|
||||
src_attn_tp_size=3,
|
||||
dst_attn_tp_size=2,
|
||||
local_tp_rank_in_group=0,
|
||||
)
|
||||
|
||||
|
||||
class TestMooncakeTransferInfoIsDummy(unittest.TestCase):
|
||||
"""Truth table for mooncake's payload-inferred is_dummy, with frames built
|
||||
as KVSender sends them: kv and aux are empty iff dummy, state indices are
|
||||
|
||||
@@ -11,7 +11,7 @@ from unittest.mock import MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
|
||||
from sglang.srt.disaggregation.base.conn import KVPoll
|
||||
from sglang.srt.disaggregation.base.conn import KVPoll, StateType
|
||||
from sglang.srt.disaggregation.common.conn import CommonKVManager
|
||||
from sglang.srt.disaggregation.common.staging_handler import PrefillStagingContext
|
||||
from sglang.srt.disaggregation.common.utils import pack_int_lists
|
||||
@@ -439,6 +439,59 @@ class TestNixlKVSenderChunkPolicy(CustomTestCase):
|
||||
self.assertTrue(sender.should_send_kv_chunk(3, last_chunk=False))
|
||||
|
||||
|
||||
class TestNixlEmptyStateTransfer(CustomTestCase):
|
||||
def test_empty_pp_state_component_is_a_noop(self):
|
||||
mgr = object.__new__(NixlKVManager)
|
||||
mgr.agent = StagingFakeAgent()
|
||||
mgr.is_mla_backend = False
|
||||
mgr.pp_size = 2
|
||||
mgr.kv_args = SimpleNamespace(prefill_start_layer=0, kv_data_ptrs=[1])
|
||||
|
||||
handle = mgr._send_kvcache_generic(
|
||||
peer_name="decode",
|
||||
src_data_ptrs=[],
|
||||
dst_data_ptrs=[],
|
||||
item_lens=[],
|
||||
prefill_data_indices=np.array([3], dtype=np.int32),
|
||||
dst_data_indices=np.array([5], dtype=np.int32),
|
||||
dst_gpu_id=0,
|
||||
notif="qsa-empty",
|
||||
state_type=StateType.QSA_PENDING,
|
||||
force_flat=True,
|
||||
src_layer_ids=[],
|
||||
dst_layer_ids=[],
|
||||
)
|
||||
|
||||
self.assertIsNone(handle)
|
||||
self.assertEqual(mgr.agent.get_xfer_descs_calls, [])
|
||||
self.assertEqual(mgr.agent.initialize_xfer_calls, [])
|
||||
|
||||
def test_paired_state_entries_reject_item_length_mismatch(self):
|
||||
mgr = object.__new__(NixlKVManager)
|
||||
mgr.agent = StagingFakeAgent()
|
||||
mgr.is_mla_backend = False
|
||||
mgr.pp_size = 1
|
||||
mgr.kv_args = SimpleNamespace(prefill_start_layer=0, kv_data_ptrs=[1])
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "item length mismatch"):
|
||||
mgr._send_kvcache_generic(
|
||||
peer_name="decode",
|
||||
src_data_ptrs=[10],
|
||||
dst_data_ptrs=[20],
|
||||
item_lens=[32],
|
||||
prefill_data_indices=np.array([3], dtype=np.int32),
|
||||
dst_data_indices=np.array([5], dtype=np.int32),
|
||||
dst_gpu_id=0,
|
||||
notif="qsa-mismatch",
|
||||
state_type=StateType.QSA_PENDING,
|
||||
force_flat=True,
|
||||
src_layer_ids=[24],
|
||||
dst_layer_ids=[24],
|
||||
dst_item_lens=[48],
|
||||
)
|
||||
self.assertEqual(mgr.agent.initialize_xfer_calls, [])
|
||||
|
||||
|
||||
class TestNixlAbortHandling(CustomTestCase):
|
||||
def _make_manager(self, request_status=None):
|
||||
mgr = object.__new__(NixlKVManager)
|
||||
|
||||
@@ -15,6 +15,8 @@ def _pool(temporal: torch.Tensor, num_conv: int = 2) -> MambaPool:
|
||||
"""A MambaPool stub carrying only what the transfer accessors read."""
|
||||
pool = object.__new__(MambaPool)
|
||||
pool.num_mamba_layers = NUM_LAYERS
|
||||
pool.mamba_layer_ids = list(range(NUM_LAYERS))
|
||||
pool._slot_siblings = []
|
||||
pool.conv_slice_axis = 0
|
||||
pool.mamba_cache = MambaPool.State(
|
||||
conv=[torch.zeros(NUM_LAYERS, NUM_SLOTS, 4, 5) for _ in range(num_conv)],
|
||||
@@ -54,6 +56,18 @@ class TestMambaStateTransferBuffers(unittest.TestCase):
|
||||
|
||||
self.assertEqual(len(pool.get_state_dim_per_tensor()), len(lens))
|
||||
|
||||
def test_sibling_declares_replicated_transfer_without_field_name_coupling(self):
|
||||
pool = _pool(torch.zeros(NUM_LAYERS, NUM_SLOTS, 6, 7, 8))
|
||||
|
||||
class ReplicatedSibling:
|
||||
def iter_transfer_state_entries(self):
|
||||
yield "future_sibling", torch.zeros(NUM_SLOTS, 9), None, 123
|
||||
|
||||
pool._slot_siblings = [ReplicatedSibling()]
|
||||
|
||||
self.assertEqual(pool.get_state_dim_per_tensor()[-1], 0)
|
||||
self.assertEqual(pool.get_state_slice_outer_counts()[-1], 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
import sys
|
||||
from contextlib import contextmanager
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
|
||||
from sglang.srt.mem_cache.qsa_kv_pool import QSATokenToKVPool
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def test_qsa_allocations_follow_parent_mooncake_scope(monkeypatch):
|
||||
active_scopes = set()
|
||||
allocations = 0
|
||||
original_zeros = torch.zeros
|
||||
|
||||
@contextmanager
|
||||
def scope(name):
|
||||
active_scopes.add(name)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
active_scopes.remove(name)
|
||||
|
||||
def init_parent(pool, **_):
|
||||
pool.full_kv_pool = SimpleNamespace(
|
||||
memory_saver_adapter=SimpleNamespace(
|
||||
region=lambda _: scope("memory_saver")
|
||||
),
|
||||
enable_custom_mem_pool=True,
|
||||
custom_mem_pool=object(),
|
||||
)
|
||||
|
||||
def allocate(*args, **kwargs):
|
||||
nonlocal allocations
|
||||
assert active_scopes == {"memory_saver", "custom_pool"}
|
||||
allocations += 1
|
||||
return original_zeros(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(HybridLinearKVPool, "__init__", init_parent)
|
||||
monkeypatch.setattr(QSATokenToKVPool, "get_kv_size_bytes", lambda _: (0, 0))
|
||||
monkeypatch.setattr(torch.cuda, "use_mem_pool", lambda _: scope("custom_pool"))
|
||||
monkeypatch.setattr("sglang.srt.mem_cache.qsa_kv_pool.torch.zeros", allocate)
|
||||
|
||||
QSATokenToKVPool(
|
||||
size=8,
|
||||
dtype=torch.bfloat16,
|
||||
page_size=4,
|
||||
head_num=1,
|
||||
head_dim=8,
|
||||
full_attention_layer_ids=[1, 3],
|
||||
device="cpu",
|
||||
mamba_pool=object(),
|
||||
qsa_index_kv_heads=1,
|
||||
qsa_index_head_dim=8,
|
||||
qsa_compress_ratio=2,
|
||||
qsa_token_topk=4,
|
||||
num_request_slots=3,
|
||||
)
|
||||
|
||||
assert allocations == 4
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
@@ -615,16 +615,22 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
# value: readers only ever read flags.
|
||||
self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "auto")
|
||||
|
||||
def test_qwen4_rejects_pd_and_unified_memory(self):
|
||||
def test_qwen4_pd_support_and_remaining_limits(self):
|
||||
qwen4 = ("Qwen4ExpForConditionalGeneration", "qwen4_exp")
|
||||
for kwargs, message in (
|
||||
({"disaggregation_mode": "prefill"}, "PD disaggregation"),
|
||||
({"disaggregation_mode": "decode"}, "PD disaggregation"),
|
||||
({"enable_unified_memory": True}, "enable-unified-memory"),
|
||||
):
|
||||
with self.subTest(**kwargs):
|
||||
with self.assertRaisesRegex(ValueError, message):
|
||||
self._construct(*qwen4, **kwargs)
|
||||
with override_platform(is_cuda=True):
|
||||
for mode in ("prefill", "decode"):
|
||||
with self.subTest(mode=mode):
|
||||
self._construct(*qwen4, disaggregation_mode=mode)
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "enable-unified-memory"):
|
||||
self._construct(*qwen4, enable_unified_memory=True)
|
||||
with self.assertRaisesRegex(ValueError, "MORI requires --pp-size 1"):
|
||||
self._construct(
|
||||
*qwen4,
|
||||
disaggregation_mode="prefill",
|
||||
disaggregation_transfer_backend="mori",
|
||||
pp_size=2,
|
||||
)
|
||||
|
||||
def test_qwen4_ple_offload_default(self):
|
||||
qwen4 = ("Qwen4ExpForConditionalGeneration", "qwen4_exp")
|
||||
|
||||
Reference in New Issue
Block a user