[2/n] [CP] Add context parallel strategy abstractions (#27313)
This commit is contained in:
@@ -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 = [
|
||||
(
|
||||
|
||||
Reference in New Issue
Block a user