[CP V1 Deprecation 3/5] Remove generic prefill CP v1 runtime (#36228)

This commit is contained in:
Baizhou Zhang
2026-09-06 21:53:53 -07:00
committed by GitHub
parent aaf9a95763
commit b6c31b155c
34 changed files with 400 additions and 741 deletions
+1 -11
View File
@@ -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(