[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)
@@ -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"]))
+15 -9
View File
@@ -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")