[3/N][CP] Implement zigzag CP strategy (#28421)

This commit is contained in:
Baizhou Zhang
2026-06-18 15:10:30 -07:00
committed by GitHub
parent 9fc9d37f6d
commit e3026ef016
13 changed files with 1091 additions and 56 deletions
+386 -1
View File
@@ -2,6 +2,8 @@ import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.layers.cp.base import (
ContextParallelStrategyKind,
get_cp_strategy,
@@ -11,10 +13,31 @@ from sglang.srt.layers.cp.base import (
is_interleave,
is_zigzag,
)
from sglang.srt.layers.cp.utils import (
cp_split_before_forward,
enable_cp_v2,
is_cp_v2_active,
)
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
from sglang.srt.runtime_context import get_parallel
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")
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class _ExtendMode:
def is_context_parallel_extend(self):
return True
class _FakeCPGroup:
def __init__(self, all_rank_tensors):
self.all_rank_tensors = all_rank_tensors
def cp_all_gather_into_tensor_async(self, output, input_tensor, stream):
del input_tensor, stream
torch.cat(self.all_rank_tensors, dim=0, out=output)
class TestCPStrategyUnit(CustomTestCase):
@@ -70,5 +93,367 @@ class TestCPStrategyUnit(CustomTestCase):
self.assertIsNotNone(get_cp_strategy())
class TestCPZigzagStrategy(CustomTestCase):
def setUp(self):
init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True,
cp_strategy="zigzag",
attn_cp_size=4,
attention_backend="fa3",
)
)
def tearDown(self):
init_cp_strategy(SimpleNamespace(enable_prefill_cp=False))
def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens):
strategy = ZigzagCPStrategy(cp_size=cp_size)
with get_parallel().override(attn_cp_rank=rank):
return strategy.build_metadata(
num_tokens=sum(extend_seq_lens),
seqs_len=seq_lens,
extend_seqs_len=extend_seq_lens,
)
def _forward_batch(self, metadata, extend_seq_lens):
return SimpleNamespace(
input_ids=torch.arange(sum(extend_seq_lens)),
forward_mode=_ExtendMode(),
extend_seq_lens_cpu=extend_seq_lens,
attn_cp_metadata=metadata,
)
def test_enable_cp_v2_and_is_cp_v2_active(self):
active_batch = SimpleNamespace(
input_ids=torch.arange(8),
forward_mode=_ExtendMode(),
extend_seq_lens_cpu=[8],
)
inactive_batch = SimpleNamespace(
input_ids=torch.arange(7),
forward_mode=_ExtendMode(),
extend_seq_lens_cpu=[7],
)
with patch(
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", return_value=False
):
self.assertFalse(enable_cp_v2())
self.assertFalse(is_cp_v2_active(active_batch))
with patch(
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", return_value=True
):
self.assertTrue(enable_cp_v2())
self.assertTrue(is_cp_v2_active(active_batch))
self.assertFalse(is_cp_v2_active(inactive_batch))
def _expected_metadata(self, *, rank, cp_size, seq_lens, extend_seq_lens):
bs = len(extend_seq_lens)
cp_segment_num = cp_size * 2
prefix_offsets = [
max(int(seq_lens[i]) - int(extend_seq_lens[i]), 0) for i in range(bs)
]
per_seq_block_sizes = []
split_list = []
for length in extend_seq_lens:
base = length // cp_segment_num
rem = length % cp_segment_num
block_sizes = [
base + 1 if block_id < rem else base
for block_id in range(cp_segment_num)
]
per_seq_block_sizes.append(block_sizes)
split_list.extend(block_sizes)
per_rank_actual_token = [
sum(
block_sizes[rank_id] + block_sizes[cp_segment_num - 1 - rank_id]
for block_sizes in per_seq_block_sizes
)
for rank_id in range(cp_size)
]
max_rank_len = [max(per_rank_actual_token)] * cp_size
zigzag_index = list(range(rank, rank + bs * cp_segment_num, cp_segment_num))
zigzag_index += list(
range(cp_segment_num - rank - 1, bs * cp_segment_num, cp_segment_num)
)
cp_reverse_index = []
for batch_id in range(bs):
cp_reverse_index.extend(
list(range(batch_id, cp_segment_num * bs, 2 * bs))
+ list(range((cp_segment_num - 1) * bs + batch_id, 0, -2 * bs))
)
reverse_split_len = []
for rank_id in range(cp_size):
for batch_id in range(bs):
reverse_split_len.append(per_seq_block_sizes[batch_id][rank_id])
for batch_id in range(bs):
reverse_split_len.append(
per_seq_block_sizes[batch_id][cp_segment_num - 1 - rank_id]
)
kv_len_prev_list = []
kv_len_next_list = []
actual_seq_q_prev_list = []
actual_seq_q_next_list = []
for batch_id, block_sizes in enumerate(per_seq_block_sizes):
kv_len_prev_list.append(
prefix_offsets[batch_id] + sum(block_sizes[: rank + 1])
)
kv_len_next_list.append(
prefix_offsets[batch_id] + sum(block_sizes[: cp_segment_num - rank])
)
actual_seq_q_prev_list.append(block_sizes[rank])
actual_seq_q_next_list.append(block_sizes[cp_segment_num - rank - 1])
return {
"bs": bs,
"total_seq_lens": sum(extend_seq_lens),
"split_list": split_list,
"zigzag_index": zigzag_index,
"per_rank_actual_token": per_rank_actual_token,
"max_rank_len": max_rank_len,
"reverse_split_len": reverse_split_len,
"cp_reverse_index": cp_reverse_index,
"kv_len_prev_list": kv_len_prev_list,
"kv_len_next_list": kv_len_next_list,
"actual_seq_q_prev_list": actual_seq_q_prev_list,
"actual_seq_q_next_list": actual_seq_q_next_list,
}
def _assert_metadata_matches(self, metadata, expected):
self.assertEqual(metadata.bs, expected["bs"])
self.assertEqual(metadata.total_seq_lens, expected["total_seq_lens"])
self.assertEqual(metadata.split_list, expected["split_list"])
self.assertEqual(metadata.zigzag_index, expected["zigzag_index"])
self.assertEqual(
metadata.per_rank_actual_token, expected["per_rank_actual_token"]
)
self.assertEqual(metadata.max_rank_len, expected["max_rank_len"])
self.assertEqual(metadata.reverse_split_len, expected["reverse_split_len"])
self.assertEqual(metadata.cp_reverse_index, expected["cp_reverse_index"])
self.assertEqual(metadata.kv_len_prev_list, expected["kv_len_prev_list"])
self.assertEqual(metadata.kv_len_next_list, expected["kv_len_next_list"])
self.assertEqual(
metadata.actual_seq_q_prev_list, expected["actual_seq_q_prev_list"]
)
self.assertEqual(
metadata.actual_seq_q_next_list, expected["actual_seq_q_next_list"]
)
self.assertEqual(
metadata.cu_seqlens_q_prev_tensor.cpu().tolist(),
[0]
+ list(
torch.tensor(expected["actual_seq_q_prev_list"]).cumsum(dim=0).tolist()
),
)
self.assertEqual(
metadata.cu_seqlens_q_next_tensor.cpu().tolist(),
[0]
+ list(
torch.tensor(expected["actual_seq_q_next_list"]).cumsum(dim=0).tolist()
),
)
def _padded_rank_tensors(self, x, *, cp_size, seq_lens, extend_seq_lens):
per_rank = []
metas = []
for rank in range(cp_size):
metadata = self._metadata_for_rank(
rank,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
metas.append(metadata)
fb = self._forward_batch(metadata, extend_seq_lens)
local = ZigzagCPStrategy(cp_size=cp_size).shard_hidden_states(x, fb)
pad = metadata.max_rank_len[0] - local.shape[0]
if pad:
local = torch.nn.functional.pad(
local,
[0, 0] * (local.ndim - 1) + [0, pad],
)
per_rank.append(local)
return metas, per_rank
def test_zigzag_metadata_for_batched_sequences(self):
cases = [
(4, [11, 13], [9, 10]),
(2, [8], [8]),
(4, [100000, 200000, 80], [100000, 200000, 64]),
(4, [100005, 200011, 25], [100000, 200000, 16]),
]
for cp_size, seq_lens, extend_seq_lens in cases:
for rank in range(cp_size):
with self.subTest(
cp_size=cp_size,
rank=rank,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
):
metadata = self._metadata_for_rank(
rank,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
expected = self._expected_metadata(
rank=rank,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
self._assert_metadata_matches(metadata, expected)
def test_zigzag_shards_hidden_states_and_position_ids(self):
cp_size = 4
seq_lens = [11, 13]
extend_seq_lens = [9, 10]
x = torch.arange(sum(extend_seq_lens) * 2).view(sum(extend_seq_lens), 2)
positions = torch.arange(sum(extend_seq_lens))
for rank in range(cp_size):
metadata = self._metadata_for_rank(
rank,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
fb = self._forward_batch(metadata, extend_seq_lens)
strategy = ZigzagCPStrategy(cp_size=cp_size)
chunks = torch.split(x, metadata.split_list, dim=0)
position_chunks = torch.split(positions, metadata.split_list, dim=-1)
expected_x = torch.cat([chunks[i] for i in metadata.zigzag_index], dim=0)
expected_positions = torch.cat(
[position_chunks[i] for i in metadata.zigzag_index], dim=-1
)
local_x = strategy.shard_hidden_states(x, fb)
local_positions = strategy.shard_position_ids(positions, fb)
with patch(
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", return_value=True
):
helper_x, helper_positions = cp_split_before_forward(
x,
positions,
fb,
)
self.assertTrue(torch.equal(local_x, expected_x))
self.assertTrue(torch.equal(local_positions, expected_positions))
self.assertTrue(torch.equal(helper_x, expected_x))
self.assertTrue(torch.equal(helper_positions, expected_positions))
def test_zigzag_gathers_hidden_states_to_original_order(self):
cp_size = 4
seq_lens = [11, 13]
extend_seq_lens = [9, 10]
x = torch.arange(sum(extend_seq_lens) * 2).view(sum(extend_seq_lens), 2)
metas, padded_rank_tensors = self._padded_rank_tensors(
x,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
for rank in range(cp_size):
local_x = padded_rank_tensors[rank][
: metas[rank].per_rank_actual_token[rank]
]
fb = self._forward_batch(metas[rank], extend_seq_lens)
with (
patch(
"sglang.srt.layers.cp.zigzag.get_attention_cp_group",
return_value=_FakeCPGroup(padded_rank_tensors),
),
patch(
"sglang.srt.distributed.device_communicators.pynccl_allocator.use_symmetric_memory",
return_value=torch.no_grad(),
),
):
gathered = ZigzagCPStrategy(cp_size=cp_size).gather_hidden_states(
local_x, fb, stream=None
)
self.assertTrue(torch.equal(gathered, x))
def test_zigzag_gathers_kv_cache_to_original_order(self):
cp_size = 4
seq_lens = [11, 13]
extend_seq_lens = [9, 10]
kv = torch.arange(sum(extend_seq_lens) * 2 * 3).view(sum(extend_seq_lens), 2, 3)
metas, padded_rank_tensors = self._padded_rank_tensors(
kv,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
for rank in range(cp_size):
local_kv = padded_rank_tensors[rank][
: metas[rank].per_rank_actual_token[rank]
]
fb = self._forward_batch(metas[rank], extend_seq_lens)
with (
patch(
"sglang.srt.layers.cp.zigzag.get_attention_cp_group",
return_value=_FakeCPGroup(padded_rank_tensors),
),
patch(
"sglang.srt.distributed.device_communicators.pynccl_allocator.use_symmetric_memory",
return_value=torch.no_grad(),
),
):
gathered = ZigzagCPStrategy(cp_size=cp_size).gather_kv_cache(
local_kv, fb, stream=None
)
self.assertTrue(torch.equal(gathered, kv))
def test_zigzag_attention_dispatch_runs_prev_then_next(self):
cp_size = 2
seq_lens = [8]
extend_seq_lens = [8]
metadata = self._metadata_for_rank(
0,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
fb = SimpleNamespace(attn_cp_metadata=metadata)
q = torch.arange(4 * 2).view(4, 2)
calls = []
def attn_fn(q_chunk, cu_seqlens_q, cache_seqlens, max_seqlen_q):
calls.append(
(
q_chunk.clone(),
cu_seqlens_q.clone(),
cache_seqlens.clone(),
max_seqlen_q,
)
)
return q_chunk + 100
out = ZigzagCPStrategy(cp_size=cp_size).run_attention(
q, fb, device=torch.device("cpu"), attn_fn=attn_fn
)
self.assertEqual(len(calls), 2)
self.assertTrue(torch.equal(calls[0][0], q[:2]))
self.assertTrue(torch.equal(calls[1][0], q[2:]))
self.assertTrue(torch.equal(out, q + 100))
if __name__ == "__main__":
unittest.main()
@@ -11,17 +11,17 @@ from sglang.test.test_utils import (
popen_launch_server,
)
register_cuda_ci(est_time=261, stage="extra-b", runner_config="4-gpu-h100")
register_cuda_ci(est_time=260, stage="extra-b", runner_config="4-gpu-h100")
QWEN3_30B_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8"
GQA_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8"
GSM8K_BASELINE_ACCURACY = 0.85
GSM8K_BASELINE_ACCURACY = 0.93
class TestQwen330B(CustomTestCase):
class TestGQACP2TP2EP2(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = QWEN3_30B_MODEL_PATH
cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
@@ -46,11 +46,13 @@ class TestQwen330B(CustomTestCase):
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
],
env={"SGLANG_ENABLE_CP_V2": "0"},
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
@@ -73,10 +75,10 @@ class TestQwen330B(CustomTestCase):
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
class TestQwen330BCP(CustomTestCase):
class TestGQACPTP2CP2EP4(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = QWEN3_30B_MODEL_PATH
cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
@@ -101,11 +103,13 @@ class TestQwen330BCP(CustomTestCase):
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
],
env={"SGLANG_ENABLE_CP_V2": "0"},
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
+200
View File
@@ -0,0 +1,200 @@
import unittest
from types import SimpleNamespace
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
kill_process_tree,
popen_launch_server,
)
register_cuda_ci(est_time=500, stage="extra-b", runner_config="4-gpu-h100")
GQA_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8"
GSM8K_BASELINE_ACCURACY = 0.93
class TestGQACP2TP2EP2(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--moe-dp-size",
"2",
"--ep-size",
"2",
"--attn-cp-size",
"2",
"--enable-prefill-cp",
"--cp-strategy",
"zigzag",
"--cuda-graph-max-bs",
"32",
"--max-running-requests",
"32",
"--trust-remote-code",
"--disable-piecewise-cuda-graph",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
],
env={"SGLANG_ENABLE_CP_V2": "1"},
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
model=self.model,
eval_name="gsm8k",
num_shots=5,
num_examples=200,
max_tokens=16000,
num_threads=128,
repeat=1,
temperature=0.6,
top_p=0.95,
top_k=20,
base_url=self.base_url,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
class TestGQACPTP2CP2EP4(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--moe-dp-size",
"1",
"--ep-size",
"4",
"--attn-cp-size",
"2",
"--enable-prefill-cp",
"--cp-strategy",
"zigzag",
"--cuda-graph-max-bs",
"32",
"--max-running-requests",
"32",
"--trust-remote-code",
"--disable-piecewise-cuda-graph",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
],
env={"SGLANG_ENABLE_CP_V2": "1"},
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
model=self.model,
eval_name="gsm8k",
num_shots=5,
num_examples=200,
max_tokens=16000,
num_threads=128,
repeat=1,
temperature=0.6,
top_p=0.95,
top_k=20,
base_url=self.base_url,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
class TestGQACPCP4EP4(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--ep",
"4",
"--attn-cp-size",
"4",
"--enable-prefill-cp",
"--cp-strategy",
"zigzag",
"--moe-a2a-backend",
"deepep",
"--attention-backend",
"fa3",
"--cuda-graph-max-bs",
"32",
"--max-running-requests",
"32",
"--trust-remote-code",
"--disable-piecewise-cuda-graph",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
],
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
model=self.model,
eval_name="gsm8k",
num_shots=5,
num_examples=200,
max_tokens=16000,
num_threads=128,
repeat=1,
temperature=0.6,
top_p=0.95,
top_k=20,
base_url=self.base_url,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
if __name__ == "__main__":
unittest.main()
@@ -159,6 +159,7 @@ class TestContextParallelServerArgs(CustomTestCase):
enable_dsa_prefill_context_parallel=False,
enable_prefill_cp=False,
cp_strategy=None,
model_path="instance://127.0.0.1:8000/dummy",
dsa_prefill_cp_mode="round-robin-split",
prefill_cp_mode="in-seq-split",
attn_cp_size=1,