Support MiMo V2.5 with zigzag context parallelism (#29972)

This commit is contained in:
Baizhou Zhang
2026-07-19 02:36:40 -07:00
committed by GitHub
parent 377c93d54e
commit 7a03d30149
17 changed files with 368 additions and 118 deletions
+94 -2
View File
@@ -13,12 +13,18 @@ from sglang.srt.layers.cp.base import (
is_interleave,
is_zigzag,
)
from sglang.srt.layers.cp.padding import (
get_cp_padding_align_size,
pad_local_rows,
pad_logical_token_to_physical,
)
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.mem_cache.memory_pool import KVWriteLoc
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
@@ -35,8 +41,8 @@ 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
def all_gather_into_tensor(self, output, input_tensor):
del input_tensor
torch.cat(self.all_rank_tensors, dim=0, out=output)
@@ -418,6 +424,92 @@ class TestCPZigzagStrategy(CustomTestCase):
self.assertTrue(torch.equal(gathered, kv))
def test_zigzag_padding_aligns_local_tensors(self):
cp_size = 2
metadata = SimpleNamespace(
per_rank_actual_token=[7, 6],
per_rank_logical_token=None,
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,
),
):
align_size = get_cp_padding_align_size()
pad_logical_token_to_physical(metadata)
self.assertEqual(align_size, 2 * cp_size)
self.assertEqual(metadata.per_rank_logical_token, [7, 6])
self.assertEqual(metadata.per_rank_actual_token, [8, 8])
self.assertEqual(metadata.max_rank_len, [8, 8])
for logical_len in metadata.per_rank_logical_token:
local_hidden = torch.arange(logical_len * 2).view(logical_len, 2)
padded_hidden = pad_local_rows(local_hidden, metadata, dim=0)
physical_len = metadata.per_rank_actual_token[0]
self.assertEqual(padded_hidden.shape, (physical_len, 2))
self.assertTrue(torch.equal(padded_hidden[:logical_len], local_hidden))
self.assertTrue(
torch.equal(
padded_hidden[logical_len:],
local_hidden.new_zeros(physical_len - logical_len, 2),
)
)
def test_zigzag_materialize_full_kv_gathers_once_and_preserves_swa_location(self):
key = torch.arange(6).view(3, 2)
value = torch.arange(9).view(3, 3) + 10
local_kv = torch.cat([key, value], dim=-1)
cache_loc = torch.arange(3)
swa_loc = torch.arange(5) + 16
forward_batch = SimpleNamespace(
out_cache_loc=cache_loc,
encoder_out_cache_loc=torch.arange(3) + 32,
)
layer = SimpleNamespace(
is_cross_attention=False,
k_scale="key-scale",
v_scale="value-scale",
)
writes = []
pool = SimpleNamespace(set_kv_buffer=lambda *args: writes.append(args))
strategy = ZigzagCPStrategy(cp_size=2)
with (
patch.object(strategy, "gather_kv_cache", return_value=local_kv) as gather,
patch(
"sglang.srt.layers.cp.zigzag.get_token_to_kv_pool",
return_value=pool,
),
):
strategy.materialize_full_kv(forward_batch, layer, key, value, swa_loc)
gather.assert_called_once()
gathered_kv, gathered_forward_batch = gather.call_args.args
self.assertTrue(gathered_kv.is_contiguous())
self.assertTrue(torch.equal(gathered_kv, local_kv))
self.assertIs(gathered_forward_batch, forward_batch)
self.assertEqual(len(writes), 1)
written_layer, write_loc, written_key, written_value, k_scale, v_scale = writes[
0
]
self.assertIs(written_layer, layer)
self.assertIsInstance(write_loc, KVWriteLoc)
self.assertIs(write_loc.loc, cache_loc)
self.assertTrue(torch.equal(write_loc.swa_loc, swa_loc[:3]))
self.assertTrue(written_key.is_contiguous())
self.assertTrue(written_value.is_contiguous())
self.assertTrue(torch.equal(written_key, key))
self.assertTrue(torch.equal(written_value, value))
self.assertEqual((k_scale, v_scale), ("key-scale", "value-scale"))
def test_zigzag_attention_dispatch_runs_prev_then_next(self):
cp_size = 2
seq_lens = [8]
+80
View File
@@ -0,0 +1,80 @@
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, suite="nightly-8-gpu-b200", nightly=True)
MIMO_V2_MODEL_PATH = "XiaomiMiMo/MiMo-V2.5"
GSM8K_BASELINE_ACCURACY = 0.93
class TestMiMoV2ContextParallel(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = MIMO_V2_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=[
"--trust-remote-code",
"--language-only",
"--tp",
"8",
"--attn-cp-size",
"2",
"--attention-backend",
"fa4",
"--enable-prefill-cp",
"--cp-strategy",
"zigzag",
"--moe-runner-backend",
"flashinfer_trtllm",
"--moe-dense-tp-size",
"1",
"--mem-fraction-static",
"0.8",
"--chunked-prefill-size",
"8192",
],
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
metrics = run_eval(
SimpleNamespace(
model=self.model,
eval_name="gsm8k",
api="chat",
num_shots=5,
num_examples=200,
max_tokens=4096,
num_threads=8,
repeat=1,
temperature=0.0,
top_p=1.0,
base_url=self.base_url,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
)
print(f"{metrics=}")
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
if __name__ == "__main__":
unittest.main()