[2/n] [CP] Add context parallel strategy abstractions (#27313)

This commit is contained in:
Baizhou Zhang
2026-06-16 00:20:04 -07:00
committed by GitHub
parent 6c908b3a3a
commit 77f327cb6e
9 changed files with 715 additions and 1 deletions
@@ -0,0 +1,74 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.layers.cp.base import (
ContextParallelStrategyKind,
get_cp_strategy,
get_cp_strategy_kind,
init_cp_strategy,
is_cp_enabled,
is_interleave,
is_zigzag,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
class TestCPStrategyUnit(CustomTestCase):
def tearDown(self):
init_cp_strategy(SimpleNamespace(enable_prefill_cp=False))
def test_strategy_kind_maps_cli_values(self):
self.assertEqual(ContextParallelStrategyKind.NONE.value, 0)
self.assertEqual(
ContextParallelStrategyKind.from_string("zigzag"),
ContextParallelStrategyKind.ZIGZAG,
)
self.assertEqual(
ContextParallelStrategyKind.from_string("interleave"),
ContextParallelStrategyKind.INTERLEAVE,
)
self.assertEqual(ContextParallelStrategyKind.ZIGZAG.cli_value, "zigzag")
self.assertEqual(ContextParallelStrategyKind.INTERLEAVE.cli_value, "interleave")
def test_init_cp_strategy_binds_zigzag_strategy(self):
init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True,
cp_strategy="zigzag",
attn_cp_size=4,
)
)
self.assertTrue(is_cp_enabled())
self.assertTrue(is_zigzag())
self.assertFalse(is_interleave())
self.assertEqual(get_cp_strategy_kind(), ContextParallelStrategyKind.ZIGZAG)
def test_get_cp_strategy_is_initialized_under_cp_v1_and_cp_v2(self):
init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True,
cp_strategy="interleave",
attn_cp_size=4,
)
)
with patch(
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", return_value=False
):
self.assertIsNotNone(get_cp_strategy())
self.assertTrue(is_cp_enabled())
self.assertTrue(is_interleave())
with patch(
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", return_value=True
):
self.assertIsNotNone(get_cp_strategy())
if __name__ == "__main__":
unittest.main()
@@ -8,6 +8,7 @@ from unittest.mock import MagicMock, patch
import sglang.srt.server_args as server_args_module
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
from sglang.srt.layers.cp.base import is_cp_enabled, is_interleave
from sglang.srt.model_executor.cuda_graph_config import (
Backend,
CudaGraphConfig,
@@ -232,6 +233,19 @@ class TestContextParallelServerArgs(CustomTestCase):
self.assertEqual(server_args.dsa_prefill_cp_mode, "round-robin-split")
self.assertEqual(server_args.prefill_cp_mode, "round-robin-split")
def test_context_parallel_handler_initializes_cp_strategy(self):
server_args = self._new_cp_args(
enable_prefill_cp=True,
cp_strategy="interleave",
attn_cp_size=2,
tp_size=2,
)
server_args._handle_context_parallelism()
self.assertTrue(is_cp_enabled())
self.assertTrue(is_interleave())
def test_registered_cp_legacy_args_map_to_unified_strategy(self):
cases = [
(