[4/N][CP] Support interleave strategy for cp v2 (#30482)
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com> Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
This commit is contained in:
co-authored by
Xinyuan Tong
Baizhou Zhang
parent
c192145830
commit
f46d5f25b4
@@ -1,4 +1,5 @@
|
||||
import unittest
|
||||
from contextlib import contextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
@@ -13,6 +14,7 @@ from sglang.srt.layers.cp.base import (
|
||||
is_interleave,
|
||||
is_zigzag,
|
||||
)
|
||||
from sglang.srt.layers.cp.interleave import InterleaveCPStrategy
|
||||
from sglang.srt.layers.cp.padding import (
|
||||
get_cp_padding_align_size,
|
||||
pad_local_rows,
|
||||
@@ -22,6 +24,7 @@ from sglang.srt.layers.cp.utils import (
|
||||
cp_split_before_forward,
|
||||
enable_cp_v2,
|
||||
is_cp_v2_active,
|
||||
prepare_cp_forward,
|
||||
)
|
||||
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
@@ -545,5 +548,338 @@ class TestCPZigzagStrategy(CustomTestCase):
|
||||
self.assertTrue(torch.equal(out, q + 100))
|
||||
|
||||
|
||||
class TestCPInterleaveStrategy(CustomTestCase):
|
||||
def setUp(self):
|
||||
init_cp_strategy(
|
||||
SimpleNamespace(
|
||||
enable_prefill_cp=True,
|
||||
cp_strategy="interleave",
|
||||
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 = InterleaveCPStrategy(cp_size=cp_size)
|
||||
with get_parallel().override(attn_cp_rank=rank, attn_cp_size=cp_size):
|
||||
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 _rank_tensors(self, x, *, cp_size, seq_lens, extend_seq_lens):
|
||||
per_rank = []
|
||||
metas = []
|
||||
with self._patch_legacy_round_robin_mode():
|
||||
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)
|
||||
strategy = InterleaveCPStrategy(cp_size=cp_size)
|
||||
with get_parallel().override(attn_cp_rank=rank, attn_cp_size=cp_size):
|
||||
per_rank.append(strategy.shard_hidden_states(x, fb))
|
||||
return metas, per_rank
|
||||
|
||||
@contextmanager
|
||||
def _patch_legacy_round_robin_mode(self):
|
||||
with patch(
|
||||
"sglang.srt.layers.attention.dsa.utils.is_dsa_prefill_cp_round_robin_split",
|
||||
return_value=True,
|
||||
):
|
||||
yield
|
||||
|
||||
@contextmanager
|
||||
def _patch_interleave_all_gather(self, rank_tensors):
|
||||
def all_gather(output, input_tensor):
|
||||
del input_tensor
|
||||
torch.cat(rank_tensors, dim=0, out=output)
|
||||
|
||||
patchers = (
|
||||
patch(
|
||||
"sglang.srt.layers.cp.interleave.attn_cp_all_gather_into_tensor",
|
||||
side_effect=all_gather,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.cp.interleave.is_allocation_symmetric",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.cp.interleave.use_symmetric_memory",
|
||||
return_value=torch.no_grad(),
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.attention.dsa.utils.is_dsa_prefill_cp_round_robin_split",
|
||||
return_value=True,
|
||||
),
|
||||
)
|
||||
with (
|
||||
patchers[0],
|
||||
patchers[1],
|
||||
patchers[2],
|
||||
patchers[3],
|
||||
get_parallel().override(attn_cp_group=object()),
|
||||
):
|
||||
yield
|
||||
|
||||
def test_interleave_metadata_supports_shared_padding(self):
|
||||
metadata = InterleaveCPStrategy(cp_size=4).build_metadata(
|
||||
num_tokens=10,
|
||||
seqs_len=[10],
|
||||
extend_seqs_len=[10],
|
||||
)
|
||||
|
||||
self.assertEqual(metadata.per_rank_actual_token, [3, 3, 2, 2])
|
||||
with patch(
|
||||
"sglang.srt.layers.cp.padding.get_cp_padding_align_size",
|
||||
return_value=4,
|
||||
):
|
||||
pad_logical_token_to_physical(metadata)
|
||||
|
||||
self.assertEqual(metadata.per_rank_logical_token, [3, 3, 2, 2])
|
||||
self.assertEqual(metadata.per_rank_actual_token, [4, 4, 4, 4])
|
||||
self.assertEqual(metadata.max_rank_len, [4, 4, 4, 4])
|
||||
|
||||
def test_prepare_cp_forward_sizes_gather_buffer_for_all_cp_ranks(self):
|
||||
forward_batch = SimpleNamespace(
|
||||
input_ids=torch.arange(10),
|
||||
positions=torch.arange(10),
|
||||
forward_mode=_ExtendMode(),
|
||||
seq_lens_cpu=[10],
|
||||
extend_seq_lens_cpu=[10],
|
||||
attn_cp_metadata=None,
|
||||
global_num_tokens_cpu=[10],
|
||||
out_cache_loc=None,
|
||||
)
|
||||
|
||||
with (
|
||||
get_parallel().override(attn_cp_rank=2, attn_cp_size=4),
|
||||
patch(
|
||||
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get",
|
||||
return_value=True,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.cp.padding.get_cp_padding_align_size",
|
||||
return_value=4,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.dp_attention.set_local_dp_buffer_len"
|
||||
) as set_buffer_len,
|
||||
):
|
||||
prepare_cp_forward(forward_batch)
|
||||
|
||||
self.assertEqual(
|
||||
forward_batch.attn_cp_metadata.per_rank_actual_token,
|
||||
[4, 4, 4, 4],
|
||||
)
|
||||
set_buffer_len.assert_called_once_with(16)
|
||||
|
||||
def test_interleave_shards_hidden_states_and_position_ids(self):
|
||||
cp_size = 4
|
||||
seq_lens = [8]
|
||||
extend_seq_lens = [8]
|
||||
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 = InterleaveCPStrategy(cp_size=cp_size)
|
||||
expected_x = x[rank::cp_size]
|
||||
expected_positions = positions[rank::cp_size]
|
||||
|
||||
with (
|
||||
get_parallel().override(
|
||||
attn_cp_rank=rank,
|
||||
attn_cp_size=cp_size,
|
||||
),
|
||||
self._patch_legacy_round_robin_mode(),
|
||||
):
|
||||
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_interleave_padding_preserves_shard_and_gather(self):
|
||||
cp_size = 4
|
||||
total_tokens = 10
|
||||
x = torch.arange(total_tokens * 2).view(total_tokens, 2)
|
||||
rank_tensors = []
|
||||
metas = []
|
||||
|
||||
for rank in range(cp_size):
|
||||
metadata = self._metadata_for_rank(
|
||||
rank,
|
||||
cp_size=cp_size,
|
||||
seq_lens=[total_tokens],
|
||||
extend_seq_lens=[total_tokens],
|
||||
)
|
||||
with patch(
|
||||
"sglang.srt.layers.cp.padding.get_cp_padding_align_size",
|
||||
return_value=4,
|
||||
):
|
||||
pad_logical_token_to_physical(metadata)
|
||||
fb = self._forward_batch(metadata, [total_tokens])
|
||||
with get_parallel().override(attn_cp_rank=rank, attn_cp_size=cp_size):
|
||||
local_x = InterleaveCPStrategy(cp_size=cp_size).shard_hidden_states(
|
||||
x, fb
|
||||
)
|
||||
|
||||
logical_len = metadata.per_rank_logical_token[rank]
|
||||
self.assertEqual(local_x.shape[0], 4)
|
||||
self.assertTrue(torch.equal(local_x[:logical_len], x[rank::cp_size]))
|
||||
self.assertTrue(torch.count_nonzero(local_x[logical_len:]) == 0)
|
||||
metas.append(metadata)
|
||||
rank_tensors.append(local_x)
|
||||
|
||||
for rank in range(cp_size):
|
||||
fb = self._forward_batch(metas[rank], [total_tokens])
|
||||
with (
|
||||
get_parallel().override(
|
||||
attn_cp_rank=rank,
|
||||
attn_cp_size=cp_size,
|
||||
),
|
||||
self._patch_interleave_all_gather(rank_tensors),
|
||||
):
|
||||
gathered = InterleaveCPStrategy(cp_size=cp_size).gather_hidden_states(
|
||||
rank_tensors[rank], fb, stream=None
|
||||
)
|
||||
|
||||
self.assertTrue(torch.equal(gathered, x))
|
||||
|
||||
def test_interleave_gathers_hidden_states_to_original_order(self):
|
||||
cp_size = 4
|
||||
seq_lens = [10]
|
||||
extend_seq_lens = [10]
|
||||
x = torch.arange(sum(extend_seq_lens) * 2).view(sum(extend_seq_lens), 2)
|
||||
metas, rank_tensors = self._rank_tensors(
|
||||
x,
|
||||
cp_size=cp_size,
|
||||
seq_lens=seq_lens,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
)
|
||||
max_rank_len = max(t.shape[0] for t in rank_tensors)
|
||||
padded_rank_tensors = []
|
||||
for tensor in rank_tensors:
|
||||
if tensor.shape[0] < max_rank_len:
|
||||
padded = tensor.new_zeros((max_rank_len, *tensor.shape[1:]))
|
||||
padded[: tensor.shape[0]] = tensor
|
||||
padded_rank_tensors.append(padded)
|
||||
else:
|
||||
padded_rank_tensors.append(tensor)
|
||||
|
||||
for rank in range(cp_size):
|
||||
fb = self._forward_batch(metas[rank], extend_seq_lens)
|
||||
with (
|
||||
get_parallel().override(
|
||||
attn_cp_rank=rank,
|
||||
attn_cp_size=cp_size,
|
||||
),
|
||||
self._patch_interleave_all_gather(padded_rank_tensors),
|
||||
):
|
||||
gathered = InterleaveCPStrategy(cp_size=cp_size).gather_hidden_states(
|
||||
rank_tensors[rank], fb, stream=None
|
||||
)
|
||||
|
||||
self.assertTrue(torch.equal(gathered, x))
|
||||
|
||||
def test_interleave_gathers_kv_cache_to_original_order(self):
|
||||
cp_size = 4
|
||||
seq_lens = [8]
|
||||
extend_seq_lens = [8]
|
||||
kv = torch.arange(sum(extend_seq_lens) * 2 * 3).view(sum(extend_seq_lens), 2, 3)
|
||||
metas, rank_tensors = self._rank_tensors(
|
||||
kv,
|
||||
cp_size=cp_size,
|
||||
seq_lens=seq_lens,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
)
|
||||
|
||||
for rank in range(cp_size):
|
||||
fb = self._forward_batch(metas[rank], extend_seq_lens)
|
||||
with (
|
||||
get_parallel().override(
|
||||
attn_cp_rank=rank,
|
||||
attn_cp_size=cp_size,
|
||||
),
|
||||
self._patch_interleave_all_gather(rank_tensors),
|
||||
):
|
||||
gathered = InterleaveCPStrategy(cp_size=cp_size).gather_kv_cache(
|
||||
rank_tensors[rank], fb, stream=None
|
||||
)
|
||||
|
||||
self.assertTrue(torch.equal(gathered, kv))
|
||||
|
||||
def test_interleave_materializes_full_mla_kv(self):
|
||||
strategy = InterleaveCPStrategy(cp_size=2)
|
||||
forward_batch = SimpleNamespace()
|
||||
layer = object()
|
||||
k_nope = torch.arange(6).view(2, 1, 3)
|
||||
k_rope = torch.arange(4).view(2, 1, 2) + 10
|
||||
full_latent = torch.arange(20).view(4, 5)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
strategy,
|
||||
"gather_kv_cache",
|
||||
return_value=full_latent,
|
||||
) as gather,
|
||||
patch(
|
||||
"sglang.srt.layers.cp.interleave.torch.cuda.current_stream",
|
||||
return_value=None,
|
||||
),
|
||||
):
|
||||
full_k_nope, full_k_rope = strategy.materialize_full_mla_kv(
|
||||
forward_batch,
|
||||
layer,
|
||||
k_nope,
|
||||
k_rope,
|
||||
)
|
||||
|
||||
gather.assert_called_once()
|
||||
packed_kv, gathered_forward_batch, stream = gather.call_args.args
|
||||
self.assertTrue(
|
||||
torch.equal(packed_kv, torch.cat([k_nope, k_rope], dim=-1).squeeze(1))
|
||||
)
|
||||
self.assertIs(gathered_forward_batch, forward_batch)
|
||||
self.assertIsNone(stream)
|
||||
self.assertTrue(torch.equal(full_k_nope, full_latent[:, :3].unsqueeze(1)))
|
||||
self.assertTrue(torch.equal(full_k_rope, full_latent[:, 3:].unsqueeze(1)))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
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,
|
||||
is_in_ci,
|
||||
popen_launch_server,
|
||||
write_github_step_summary,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=320, stage="extra-b", runner_config="deepep-8-gpu-h200")
|
||||
GLM52_MODEL_PATH = "zai-org/GLM-5.2-FP8"
|
||||
SERVER_LAUNCH_TIMEOUT = max(DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, 1800)
|
||||
|
||||
|
||||
class TestDSACPV2Interleave(CustomTestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.model = GLM52_MODEL_PATH
|
||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||
other_args = [
|
||||
"--trust-remote-code",
|
||||
"--tp",
|
||||
"8",
|
||||
"--enable-prefill-cp",
|
||||
"--cp-strategy",
|
||||
"interleave",
|
||||
"--attn-cp-size",
|
||||
"8",
|
||||
"--speculative-algorithm",
|
||||
"EAGLE",
|
||||
"--speculative-num-steps",
|
||||
"3",
|
||||
"--speculative-eagle-topk",
|
||||
"1",
|
||||
"--speculative-num-draft-tokens",
|
||||
"4",
|
||||
"--mem-frac",
|
||||
"0.85",
|
||||
"--cuda-graph-max-bs",
|
||||
"32",
|
||||
"--max-running-requests",
|
||||
"32",
|
||||
"--model-loader-extra-config",
|
||||
'{"enable_multithread_load": true, "num_threads": 64}',
|
||||
]
|
||||
cls.process = popen_launch_server(
|
||||
cls.model,
|
||||
cls.base_url,
|
||||
timeout=SERVER_LAUNCH_TIMEOUT,
|
||||
other_args=other_args,
|
||||
env={"SGLANG_ENABLE_CP_V2": "1"},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
kill_process_tree(cls.process.pid)
|
||||
|
||||
def test_a_gsm8k(
|
||||
self,
|
||||
): # Append an "a" to make this test run first (alphabetically) to warm up the server
|
||||
args = SimpleNamespace(
|
||||
base_url=self.base_url,
|
||||
model=self.model,
|
||||
eval_name="gsm8k",
|
||||
api="completion",
|
||||
max_tokens=512,
|
||||
num_examples=500,
|
||||
num_threads=32,
|
||||
num_shots=20,
|
||||
)
|
||||
metrics = run_eval(args)
|
||||
print(f"{metrics=}")
|
||||
|
||||
if is_in_ci():
|
||||
write_github_step_summary(
|
||||
f"### test_a_gsm8k (dsa-cp-v2-interleave)\n"
|
||||
f'{metrics["score"]=:.3f}\n'
|
||||
)
|
||||
self.assertGreater(metrics["score"], 0.935)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+1
@@ -122,6 +122,7 @@ class TestDeepseekV32CPRoundRobinSplit(CustomTestCase):
|
||||
cls.base_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=other_args,
|
||||
env={"SGLANG_ENABLE_CP_V2": "0"},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
Reference in New Issue
Block a user