[Qwen3.8-Next] Add PD state transfer for Flash Next (#36651)

This commit is contained in:
YAMY
2026-09-11 22:31:45 -07:00
committed by GitHub
parent dbd4302bbb
commit 55e5e21c88
18 changed files with 734 additions and 98 deletions
@@ -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)