[bug-fix] Stabilize GLM-5.2 MTP IndexShare across PD and CUDA graph replay (#30839)
Co-authored-by: kpham-sgl <khoa.pham@radixark.ai> Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
kpham-sgl
Claude Fable 5
parent
41ad0d9c26
commit
78dc581518
@@ -261,6 +261,7 @@ class MockModelRunner:
|
||||
"dsa_decode_backend": "fa3",
|
||||
"dsa_topk_backend": "sgl-kernel",
|
||||
"dsa_paged_mqa_logits_backend": "auto",
|
||||
"disaggregation_mode": "null",
|
||||
},
|
||||
)()
|
||||
self.hisparse_coordinator = None
|
||||
|
||||
@@ -13,10 +13,15 @@ from sglang.srt.disaggregation.common.utils import (
|
||||
unpack_list_of_buffers,
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import (
|
||||
MetadataBuffers,
|
||||
get_dsv4_c128_state_indices,
|
||||
setup_state_kv_args,
|
||||
)
|
||||
from sglang.srt.managers.overlap_utils import FutureMap, RelayPayload
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.speculative.eagle_disaggregation import (
|
||||
build_eagle_disagg_draft_input,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
|
||||
@@ -93,6 +98,97 @@ class TestGroupConcurrentContiguous(unittest.TestCase):
|
||||
group_concurrent_contiguous(self._arr([1, 2, 3]), self._arr([1, 2]))
|
||||
|
||||
|
||||
class TestEagleDsaSeedTransfer(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _make_req(seed, metadata_buffer_index=0):
|
||||
return SimpleNamespace(
|
||||
metadata_buffer_index=metadata_buffer_index,
|
||||
output_ids=[101],
|
||||
cached_tokens=0,
|
||||
cached_tokens_device=0,
|
||||
cached_tokens_host=0,
|
||||
cached_tokens_storage=0,
|
||||
multimodal_inputs=None,
|
||||
return_logprob=False,
|
||||
return_sampling_mask=False,
|
||||
hidden_states_tensor=torch.tensor([1.0, 2.0]),
|
||||
output_topk_p=torch.tensor([1.0]),
|
||||
output_topk_index=torch.tensor([7]),
|
||||
output_dsa_topk_indices=seed,
|
||||
bootstrap_room=9,
|
||||
)
|
||||
|
||||
def test_metadata_buffer_copies_seed_and_uses_invalid_sentinel(self):
|
||||
buffers = MetadataBuffers(
|
||||
size=2,
|
||||
hidden_size=2,
|
||||
hidden_states_dtype=torch.float32,
|
||||
output_dsa_topk_indices_dim=3,
|
||||
)
|
||||
seed = torch.tensor([4, 5, 6], dtype=torch.int32)
|
||||
buffers.set_buf(self._make_req(seed))
|
||||
buffers.set_buf(self._make_req(None, metadata_buffer_index=1))
|
||||
|
||||
self.assertTrue(torch.equal(buffers.output_dsa_topk_indices[0], seed))
|
||||
self.assertEqual(buffers.output_dsa_topk_indices[1].tolist(), [-1, -1, -1])
|
||||
ptrs, data_lens, item_lens = buffers.get_buf_infos()
|
||||
self.assertEqual(ptrs[-2], buffers.output_dsa_topk_indices.data_ptr())
|
||||
self.assertEqual(data_lens[-2], buffers.output_dsa_topk_indices.nbytes)
|
||||
self.assertEqual(item_lens[-2], buffers.output_dsa_topk_indices[0].nbytes)
|
||||
|
||||
def test_decode_input_requires_valid_seed_for_every_request(self):
|
||||
seeds = (
|
||||
torch.tensor([1, 2, 3], dtype=torch.int32),
|
||||
torch.tensor([4, 5, 6], dtype=torch.int32),
|
||||
)
|
||||
batch = SimpleNamespace(
|
||||
reqs=[self._make_req(seed) for seed in seeds],
|
||||
device="cpu",
|
||||
enable_overlap=False,
|
||||
)
|
||||
server_args = SimpleNamespace(
|
||||
speculative_eagle_topk=1,
|
||||
speculative_num_steps=5,
|
||||
enable_multi_layer_eagle=False,
|
||||
)
|
||||
last_tokens = torch.tensor([11, 12], dtype=torch.int64)
|
||||
|
||||
draft_input = build_eagle_disagg_draft_input(
|
||||
batch, server_args, last_tokens, None
|
||||
)
|
||||
self.assertTrue(torch.equal(draft_input.dsa_topk_indices, torch.stack(seeds)))
|
||||
|
||||
for invalid_seed in (
|
||||
None,
|
||||
torch.full((3,), -1, dtype=torch.int32),
|
||||
):
|
||||
batch.reqs[1].output_dsa_topk_indices = invalid_seed
|
||||
draft_input = build_eagle_disagg_draft_input(
|
||||
batch, server_args, last_tokens, None
|
||||
)
|
||||
self.assertIsNone(draft_input.dsa_topk_indices)
|
||||
|
||||
def test_future_map_initializes_seed_buffer_after_seedless_payload(self):
|
||||
future_map = object.__new__(FutureMap)
|
||||
future_map.dsa_topk_indices_buf = None
|
||||
future_map.req_pool_size = 4
|
||||
future_map.device = "cpu"
|
||||
future_map._maybe_init_dsa_topk_indices_buf(
|
||||
RelayPayload(bonus_tokens=torch.zeros((2,), dtype=torch.int64))
|
||||
)
|
||||
self.assertIsNone(future_map.dsa_topk_indices_buf)
|
||||
|
||||
seeds = torch.tensor([[1, 2, 3], [4, 5, 6]], dtype=torch.int32)
|
||||
future_map._maybe_init_dsa_topk_indices_buf(
|
||||
RelayPayload(
|
||||
bonus_tokens=torch.zeros((2,), dtype=torch.int64),
|
||||
dsa_topk_indices=seeds,
|
||||
)
|
||||
)
|
||||
self.assertEqual(future_map.dsa_topk_indices_buf.shape, (4, 3))
|
||||
self.assertEqual(future_map.dsa_topk_indices_buf.dtype, torch.int32)
|
||||
|
||||
|
||||
class TestDSV4C128StateIndices(unittest.TestCase):
|
||||
def test_online_aligned_boundary_has_no_partial_state(self):
|
||||
np.testing.assert_array_equal(
|
||||
|
||||
@@ -8,10 +8,11 @@ slow path (`organize_draft_results`) for num_steps in {1, 2, 3, 4}.
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.speculative.adaptive_runtime_state import SpecRuntimeState
|
||||
from sglang.srt.speculative.eagle_utils import organize_draft_results
|
||||
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
|
||||
@@ -71,10 +72,11 @@ def _make_worker(num_steps: int, num_draft_tokens: int):
|
||||
return worker
|
||||
|
||||
|
||||
def _make_backend_factory(decode_backend, draft_extend_backend):
|
||||
def _make_backend_factory(decode_backend, draft_extend_backend, captured_kwargs=None):
|
||||
class FakeDraftBackendFactory:
|
||||
def __init__(self, *args, **kwargs):
|
||||
pass
|
||||
if captured_kwargs is not None:
|
||||
captured_kwargs.update(kwargs)
|
||||
|
||||
def create_decode_backend(self):
|
||||
return decode_backend
|
||||
@@ -122,6 +124,78 @@ class TestEagleWorkerV2Topk1FastPath(CustomTestCase):
|
||||
|
||||
|
||||
class TestEagleWorkerV2BackendFallback(CustomTestCase):
|
||||
def test_missing_seed_cuda_graph_fallback(self):
|
||||
graph_result = (
|
||||
[],
|
||||
torch.zeros((1, 1), dtype=torch.long, device=DEVICE),
|
||||
torch.zeros((1, 1), dtype=torch.long, device=DEVICE),
|
||||
None,
|
||||
)
|
||||
tree_result = (
|
||||
torch.empty((0,), dtype=torch.bool, device=DEVICE),
|
||||
torch.zeros((1,), dtype=torch.long, device=DEVICE),
|
||||
torch.zeros((1, 2), dtype=torch.long, device=DEVICE),
|
||||
torch.zeros((1, 2), dtype=torch.long, device=DEVICE),
|
||||
torch.zeros((1, 2), dtype=torch.long, device=DEVICE),
|
||||
torch.zeros((2,), dtype=torch.long, device=DEVICE),
|
||||
)
|
||||
|
||||
for seed_enabled, seed_present, expect_graph in (
|
||||
(True, False, False),
|
||||
(True, True, True),
|
||||
(False, False, True),
|
||||
):
|
||||
with self.subTest(
|
||||
seed_enabled=seed_enabled,
|
||||
seed_present=seed_present,
|
||||
):
|
||||
worker = object.__new__(EagleDraftWorker)
|
||||
worker.req_to_token_pool = None
|
||||
worker.cuda_graph_runner = SimpleNamespace(
|
||||
execute=MagicMock(return_value=graph_result)
|
||||
)
|
||||
worker.draft_runner = SimpleNamespace(canary_manager=None)
|
||||
worker.topk = 1
|
||||
worker.speculative_num_steps = 1
|
||||
worker.speculative_num_draft_tokens = 2
|
||||
worker.device = DEVICE
|
||||
worker.tree_mask_mode = None
|
||||
worker.seed_dsa_topk_from_draft_extend = seed_enabled
|
||||
worker.index_share_for_mtp_iteration = True
|
||||
forward_batch = SimpleNamespace(forward_mode=ForwardMode.DECODE)
|
||||
worker.prepare_for_draft = MagicMock(return_value=(forward_batch, True))
|
||||
worker.draft_forward = MagicMock(return_value=graph_result)
|
||||
attn_backend = SimpleNamespace(
|
||||
get_verify_buffers_to_fill_after_draft=lambda: (None, None),
|
||||
max_context_len=1,
|
||||
)
|
||||
worker.target_worker = SimpleNamespace(
|
||||
model_runner=SimpleNamespace(attn_backend=attn_backend)
|
||||
)
|
||||
draft_input = SimpleNamespace(
|
||||
bonus_tokens=torch.zeros((1,), dtype=torch.long, device=DEVICE),
|
||||
dsa_topk_indices=(
|
||||
torch.ones((1, 1), dtype=torch.int32, device=DEVICE)
|
||||
if seed_present
|
||||
else None
|
||||
),
|
||||
)
|
||||
batch = SimpleNamespace(
|
||||
spec_info=draft_input,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
seq_lens_sum=1,
|
||||
seq_lens=torch.ones((1,), dtype=torch.int32, device=DEVICE),
|
||||
)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.speculative.eagle_worker_v2.build_tree_kernel_efficient",
|
||||
return_value=tree_result,
|
||||
):
|
||||
worker.draft(batch)
|
||||
|
||||
self.assertEqual(worker.cuda_graph_runner.execute.called, expect_graph)
|
||||
self.assertEqual(worker.draft_forward.called, not expect_graph)
|
||||
|
||||
def test_preserves_initialized_backend_when_draft_extend_backend_is_unset(self):
|
||||
worker = object.__new__(EagleDraftWorker)
|
||||
existing_backend = object()
|
||||
@@ -130,6 +204,7 @@ class TestEagleWorkerV2BackendFallback(CustomTestCase):
|
||||
worker.draft_runner = SimpleNamespace(attn_backend=existing_backend)
|
||||
worker.topk = 1
|
||||
worker.speculative_num_steps = 2
|
||||
worker.seed_dsa_topk_from_draft_extend = False
|
||||
|
||||
with patch(
|
||||
"sglang.srt.speculative.eagle_worker_v2.DraftBackendFactory",
|
||||
@@ -151,10 +226,14 @@ class TestEagleWorkerV2BackendFallback(CustomTestCase):
|
||||
worker.draft_runner = SimpleNamespace(attn_backend=existing_backend)
|
||||
worker.topk = 1
|
||||
worker.speculative_num_steps = 2
|
||||
worker.seed_dsa_topk_from_draft_extend = True
|
||||
factory_kwargs = {}
|
||||
|
||||
with patch(
|
||||
"sglang.srt.speculative.eagle_worker_v2.DraftBackendFactory",
|
||||
_make_backend_factory(decode_backend, draft_extend_backend),
|
||||
_make_backend_factory(
|
||||
decode_backend, draft_extend_backend, captured_kwargs=factory_kwargs
|
||||
),
|
||||
):
|
||||
worker.init_attention_backend()
|
||||
|
||||
@@ -162,6 +241,7 @@ class TestEagleWorkerV2BackendFallback(CustomTestCase):
|
||||
self.assertIs(worker.draft_extend_attn_backend, draft_extend_backend)
|
||||
self.assertIs(worker.draft_runner.draft_attn_backend, decode_backend)
|
||||
self.assertIs(worker.draft_runner.attn_backend, draft_extend_backend)
|
||||
self.assertTrue(factory_kwargs["seed_dsa_topk_from_draft_extend"])
|
||||
|
||||
def _make_adaptive_worker(self, runner_attn_backend):
|
||||
"""An EAGLEWorkerV2 with a draft worker whose state-machine fields are
|
||||
|
||||
Reference in New Issue
Block a user