[CP V1 Deprecation 3/5] Remove generic prefill CP v1 runtime (#36228)
This commit is contained in:
@@ -597,17 +597,7 @@ class TestCPZigzagStrategy(CustomTestCase):
|
||||
max_rank_len=[7, 7],
|
||||
)
|
||||
|
||||
with (
|
||||
get_parallel().override(attn_cp_size=cp_size),
|
||||
patch(
|
||||
"sglang.srt.layers.utils.cp_utils.is_prefill_cp_in_seq_split",
|
||||
return_value=True,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.attention.dsa.utils.is_dsa_prefill_cp_in_seq_split",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
with get_parallel().override(attn_cp_size=cp_size):
|
||||
align_size = get_cp_padding_align_size()
|
||||
pad_logical_token_to_physical(metadata)
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.index_topk_share import IndexTopKShareState
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
||||
@@ -15,9 +16,7 @@ def _batch(
|
||||
) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
reuse_dsa_topk_indices=reuse,
|
||||
forward_mode=SimpleNamespace(
|
||||
is_extend=lambda include_draft_extend_v2: is_extend
|
||||
),
|
||||
forward_mode=ForwardMode.DRAFT_EXTEND_V2 if is_extend else ForwardMode.DECODE,
|
||||
spec_info=SimpleNamespace(
|
||||
dsa_topk_indices=carried,
|
||||
dsa_seed_topk_capture=seed_buf,
|
||||
|
||||
@@ -14,6 +14,11 @@ import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
from sglang.srt.layers.cp import base as cp_base
|
||||
from sglang.srt.layers.cp import utils as cp_utils
|
||||
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
|
||||
from sglang.srt.layers.utils import cp_utils as platform_cp_utils
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.models.deepseek_common import attention_backend_handler as abh
|
||||
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import (
|
||||
AttnForwardMethod,
|
||||
@@ -95,5 +100,55 @@ class TestResolveRocmForwardMethod(CustomTestCase):
|
||||
self.assertEqual(abh.resolve_rocm_forward_method(method), method)
|
||||
|
||||
|
||||
class TestCPMLADispatch(CustomTestCase):
|
||||
def test_strategy_cp_uses_absorbed_mla_without_legacy_flags(self):
|
||||
# Normal MHA writes rank-local KV against full out_cache_loc before
|
||||
# the CP backend can gather it. Both one-shot and chunked MHA must
|
||||
# therefore be bypassed for an active strategy-based CP batch.
|
||||
attn = SimpleNamespace(
|
||||
chunked_prefix_cache_threshold=0,
|
||||
disable_chunked_prefix_cache=False,
|
||||
flashinfer_mla_disable_ragged=False,
|
||||
)
|
||||
with (
|
||||
mock.patch.object(abh, "_is_hip", False),
|
||||
mock.patch.object(cp_utils, "enable_cp_v2", return_value=True),
|
||||
mock.patch.object(cp_base, "_STRATEGY", ZigzagCPStrategy(cp_size=4)),
|
||||
mock.patch.object(
|
||||
platform_cp_utils,
|
||||
"get_parallel",
|
||||
return_value=SimpleNamespace(enable_prefill_context_parallel=False),
|
||||
),
|
||||
):
|
||||
for prefix in (0, 32):
|
||||
for capacity in (0, 8192):
|
||||
for num_tokens in (1, 3952):
|
||||
with self.subTest(
|
||||
prefix=prefix, capacity=capacity, num_tokens=num_tokens
|
||||
):
|
||||
batch = SimpleNamespace(
|
||||
forward_mode=ForwardMode.EXTEND,
|
||||
input_ids=range(num_tokens),
|
||||
attn_cp_metadata=None,
|
||||
extend_prefix_lens_cpu=[prefix],
|
||||
extend_seq_lens_cpu=[num_tokens],
|
||||
seq_lens_cpu=[prefix + num_tokens],
|
||||
get_max_chunk_capacity=lambda: capacity,
|
||||
)
|
||||
expected = (
|
||||
AttnForwardMethod.MLA
|
||||
if num_tokens == 3952
|
||||
else (
|
||||
AttnForwardMethod.MHA_ONE_SHOT
|
||||
if capacity == 8192
|
||||
else AttnForwardMethod.MHA_CHUNKED_KV
|
||||
)
|
||||
)
|
||||
self.assertEqual(
|
||||
abh._handle_attention_backend(attn, batch, "fa3"),
|
||||
expected,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -1075,28 +1075,22 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
with self.assertRaisesRegex(ValueError, "DeepSeek V3.2.*interleave"):
|
||||
handle_context_parallelism(server_args)
|
||||
|
||||
@override_platform(is_hip=False, is_npu=False)
|
||||
def test_generic_canonical_cp_mirrors_to_transitional_runtime_fields(self):
|
||||
@override_platform(is_hip=False, is_npu=False, is_musa=False)
|
||||
def test_generic_canonical_cp_does_not_enable_platform_runtime_fields(self):
|
||||
cases = (
|
||||
(
|
||||
"zigzag_mla_or_gqa",
|
||||
"zigzag",
|
||||
"fa3",
|
||||
True,
|
||||
False,
|
||||
"in-seq-split",
|
||||
),
|
||||
(
|
||||
"interleave_dsa",
|
||||
"interleave",
|
||||
"dsa",
|
||||
False,
|
||||
True,
|
||||
"round-robin-split",
|
||||
),
|
||||
)
|
||||
|
||||
for name, strategy, backend, expect_generic, expect_dsa, mode in cases:
|
||||
for name, strategy, backend in cases:
|
||||
with self.subTest(name=name):
|
||||
server_args = self._new_cp_args(
|
||||
enable_prefill_cp=True,
|
||||
@@ -1105,6 +1099,7 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
)
|
||||
|
||||
handle_platform_cp_compatibility(server_args)
|
||||
handle_legacy_cp_runtime_compatibility(server_args)
|
||||
|
||||
self.assertFalse(
|
||||
resolution_result(server_args, "enable_prefill_context_parallel")
|
||||
@@ -1115,32 +1110,13 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
)
|
||||
)
|
||||
|
||||
handle_legacy_cp_runtime_compatibility(server_args)
|
||||
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "enable_prefill_context_parallel"),
|
||||
expect_generic,
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(
|
||||
server_args, "enable_dsa_prefill_context_parallel"
|
||||
),
|
||||
expect_dsa,
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "dsa_prefill_cp_mode"), mode
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "prefill_cp_mode"), mode
|
||||
)
|
||||
|
||||
@override_platform(is_hip=False, is_npu=False)
|
||||
@override_platform(is_hip=False, is_npu=False, is_musa=False)
|
||||
def test_non_platform_legacy_prefill_cp_is_rejected(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="instance://127.0.0.1:8000/dummy",
|
||||
enable_prefill_context_parallel=True,
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "HIP or Ascend NPU"):
|
||||
with self.assertRaisesRegex(ValueError, "protected HIP, Ascend NPU, or MUSA"):
|
||||
handle_platform_cp_compatibility(server_args)
|
||||
|
||||
def test_generic_v1_cp_options_are_not_public_cli(self):
|
||||
@@ -1172,28 +1148,36 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
resolution_result(args, "dsa_prefill_cp_mode"), "round-robin-split"
|
||||
)
|
||||
|
||||
def test_canonical_interleave_cp_mirrors_to_dsa_runtime_aliases(self):
|
||||
server_args = self._new_cp_args(
|
||||
enable_prefill_cp=True,
|
||||
cp_strategy="interleave",
|
||||
attention_backend="dsa",
|
||||
)
|
||||
def test_platform_interleave_cp_mirrors_to_dsa_runtime_aliases(self):
|
||||
for platform in ("is_hip", "is_npu", "is_musa"):
|
||||
facts = dict(is_hip=False, is_npu=False, is_musa=False)
|
||||
facts[platform] = True
|
||||
with self.subTest(platform=platform), override_platform(**facts):
|
||||
server_args = self._new_cp_args(
|
||||
enable_prefill_cp=True,
|
||||
cp_strategy="interleave",
|
||||
attention_backend="dsa",
|
||||
)
|
||||
|
||||
handle_legacy_cp_runtime_compatibility(server_args)
|
||||
handle_context_parallelism(server_args)
|
||||
handle_legacy_cp_runtime_compatibility(server_args)
|
||||
handle_context_parallelism(server_args)
|
||||
|
||||
self.assertTrue(
|
||||
resolution_result(server_args, "enable_dsa_prefill_context_parallel")
|
||||
)
|
||||
self.assertFalse(
|
||||
resolution_result(server_args, "enable_prefill_context_parallel")
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "dsa_prefill_cp_mode"), "round-robin-split"
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "prefill_cp_mode"), "round-robin-split"
|
||||
)
|
||||
self.assertTrue(
|
||||
resolution_result(
|
||||
server_args, "enable_dsa_prefill_context_parallel"
|
||||
)
|
||||
)
|
||||
self.assertFalse(
|
||||
resolution_result(server_args, "enable_prefill_context_parallel")
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "dsa_prefill_cp_mode"),
|
||||
"round-robin-split",
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "prefill_cp_mode"),
|
||||
"round-robin-split",
|
||||
)
|
||||
|
||||
def test_context_parallel_handler_initializes_cp_strategy(self):
|
||||
server_args = self._new_cp_args(
|
||||
|
||||
Reference in New Issue
Block a user