Support multi sharding group on the same dimension in dump comparator (#19601)

This commit is contained in:
fzyzcjy
2026-03-01 10:36:48 +08:00
committed by GitHub
parent 46960e65cf
commit ea6ff7b01f
14 changed files with 469 additions and 151 deletions
@@ -90,7 +90,7 @@ class TestEnsureDimsInMetas:
assert result is metas
def test_cp_sharded_sglang_input_ids_infers_dims(self):
"""CP + input_ids in sglang infers dims 't(cp,zigzag)'."""
"""CP + input_ids in sglang infers dims 't(cp:zigzag)'."""
metas: list[dict] = [
self._make_meta(cp_size=2, cp_rank=0),
self._make_meta(cp_size=2, cp_rank=1),
@@ -99,11 +99,11 @@ class TestEnsureDimsInMetas:
name="input_ids", plugin=_sglang_plugin, metas=metas, ndim=1
)
assert result is not metas
assert result[0]["dims"] == "t(cp,zigzag)"
assert result[1]["dims"] == "t(cp,zigzag)"
assert result[0]["dims"] == "t(cp:zigzag)"
assert result[1]["dims"] == "t(cp:zigzag)"
def test_cp_sharded_sglang_positions_infers_dims(self):
"""CP + positions in sglang infers dims 't(cp,zigzag)'."""
"""CP + positions in sglang infers dims 't(cp:zigzag)'."""
metas: list[dict] = [
self._make_meta(cp_size=2, cp_rank=0),
self._make_meta(cp_size=2, cp_rank=1),
@@ -111,10 +111,10 @@ class TestEnsureDimsInMetas:
result = _ensure_dims_in_metas(
name="positions", plugin=_sglang_plugin, metas=metas, ndim=1
)
assert result[0]["dims"] == "t(cp,zigzag)"
assert result[0]["dims"] == "t(cp:zigzag)"
def test_cp_sharded_megatron_input_ids_infers_dims_1d(self):
"""CP + input_ids in megatron (1D) infers dims 't(cp,zigzag)'."""
"""CP + input_ids in megatron (1D) infers dims 't(cp:zigzag)'."""
metas: list[dict] = [
{"megatron_parallel_info": {"cp_rank": 0, "cp_size": 2}},
{"megatron_parallel_info": {"cp_rank": 1, "cp_size": 2}},
@@ -122,10 +122,10 @@ class TestEnsureDimsInMetas:
result = _ensure_dims_in_metas(
name="input_ids", plugin=_megatron_plugin, metas=metas, ndim=1
)
assert result[0]["dims"] == "t(cp,zigzag)"
assert result[0]["dims"] == "t(cp:zigzag)"
def test_cp_sharded_megatron_input_ids_infers_dims_2d(self):
"""CP + input_ids in megatron (2D) infers dims 'b s(cp,zigzag)'."""
"""CP + input_ids in megatron (2D) infers dims 'b s(cp:zigzag)'."""
metas: list[dict] = [
{"megatron_parallel_info": {"cp_rank": 0, "cp_size": 2}},
{"megatron_parallel_info": {"cp_rank": 1, "cp_size": 2}},
@@ -133,7 +133,7 @@ class TestEnsureDimsInMetas:
result = _ensure_dims_in_metas(
name="input_ids", plugin=_megatron_plugin, metas=metas, ndim=2
)
assert result[0]["dims"] == "b s(cp,zigzag)"
assert result[0]["dims"] == "b s(cp:zigzag)"
def test_cp_non_sharded_name_returns_metas_unchanged(self):
"""CP + non-sharded tensor name (seq_lens) returns metas as-is."""
@@ -218,19 +218,19 @@ class TestInferCpShardedDims:
"""Tests for infer_cp_sharded_dims on each plugin."""
def test_megatron_infer_1d(self) -> None:
"""Megatron 1D → 't(cp,zigzag)'."""
"""Megatron 1D → 't(cp:zigzag)'."""
result: str = _megatron_plugin.infer_cp_sharded_dims(name="input_ids", ndim=1)
assert result == "t(cp,zigzag)"
assert result == "t(cp:zigzag)"
def test_megatron_infer_2d(self) -> None:
"""Megatron 2D → 'b s(cp,zigzag)'."""
"""Megatron 2D → 'b s(cp:zigzag)'."""
result: str = _megatron_plugin.infer_cp_sharded_dims(name="input_ids", ndim=2)
assert result == "b s(cp,zigzag)"
assert result == "b s(cp:zigzag)"
def test_sglang_infer_1d(self) -> None:
"""SGLang 1D → 't(cp,zigzag)'."""
"""SGLang 1D → 't(cp:zigzag)'."""
result: str = _sglang_plugin.infer_cp_sharded_dims(name="input_ids", ndim=1)
assert result == "t(cp,zigzag)"
assert result == "t(cp:zigzag)"
def test_megatron_infer_3d_raises(self) -> None:
"""Megatron 3D raises ValueError."""