Support token align with packed CP data in dump comparator (#19463)

This commit is contained in:
fzyzcjy
2026-02-27 08:12:54 +08:00
committed by GitHub
parent 695e93b91f
commit 8293a914a6
6 changed files with 254 additions and 61 deletions
@@ -77,7 +77,7 @@ class TestEnsureDimsInMetas:
"""Without CP parallelism, metas are returned as-is."""
metas: list[dict] = [self._make_meta(cp_size=1)]
result = _ensure_dims_in_metas(
name="input_ids", plugin=_sglang_plugin, metas=metas
name="input_ids", plugin=_sglang_plugin, metas=metas, ndim=1
)
assert result is metas
@@ -85,38 +85,55 @@ class TestEnsureDimsInMetas:
"""If dims is already in meta, metas are returned as-is."""
metas: list[dict] = [{**self._make_meta(cp_size=2, cp_rank=0), "dims": "t"}]
result = _ensure_dims_in_metas(
name="input_ids", plugin=_sglang_plugin, metas=metas
name="input_ids", plugin=_sglang_plugin, metas=metas, ndim=1
)
assert result is metas
def test_cp_sharded_sglang_input_ids_raises(self):
"""CP + input_ids in sglang raises NotImplementedError."""
def test_cp_sharded_sglang_input_ids_infers_dims(self):
"""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),
]
with pytest.raises(NotImplementedError, match="CP-sharded"):
_ensure_dims_in_metas(name="input_ids", plugin=_sglang_plugin, metas=metas)
result = _ensure_dims_in_metas(
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)"
def test_cp_sharded_sglang_positions_raises(self):
"""CP + positions in sglang raises NotImplementedError."""
def test_cp_sharded_sglang_positions_infers_dims(self):
"""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),
]
with pytest.raises(NotImplementedError, match="CP-sharded"):
_ensure_dims_in_metas(name="positions", plugin=_sglang_plugin, metas=metas)
result = _ensure_dims_in_metas(
name="positions", plugin=_sglang_plugin, metas=metas, ndim=1
)
assert result[0]["dims"] == "t(cp,zigzag)"
def test_cp_sharded_megatron_input_ids_raises(self):
"""CP + input_ids in megatron raises NotImplementedError."""
def test_cp_sharded_megatron_input_ids_infers_dims_1d(self):
"""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}},
]
with pytest.raises(NotImplementedError, match="CP-sharded"):
_ensure_dims_in_metas(
name="input_ids", plugin=_megatron_plugin, metas=metas
)
result = _ensure_dims_in_metas(
name="input_ids", plugin=_megatron_plugin, metas=metas, ndim=1
)
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)'."""
metas: list[dict] = [
{"megatron_parallel_info": {"cp_rank": 0, "cp_size": 2}},
{"megatron_parallel_info": {"cp_rank": 1, "cp_size": 2}},
]
result = _ensure_dims_in_metas(
name="input_ids", plugin=_megatron_plugin, metas=metas, ndim=2
)
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."""
@@ -125,7 +142,7 @@ class TestEnsureDimsInMetas:
self._make_meta(cp_size=2, cp_rank=1),
]
result = _ensure_dims_in_metas(
name="seq_lens", plugin=_sglang_plugin, metas=metas
name="seq_lens", plugin=_sglang_plugin, metas=metas, ndim=1
)
assert result is metas
@@ -142,7 +159,7 @@ class TestEnsureDimsInMetas:
self._make_meta(cp_size=2, cp_rank=1),
]
result = _ensure_dims_in_metas(
name="input_ids", plugin=_DummyPlugin(), metas=metas
name="input_ids", plugin=_DummyPlugin(), metas=metas, ndim=1
)
assert result is metas
@@ -214,5 +214,34 @@ class TestInferPositions:
assert torch.equal(result, torch.tensor([0, 1, 0, 1, 2]))
class TestInferCpShardedDims:
"""Tests for infer_cp_sharded_dims on each plugin."""
def test_megatron_infer_1d(self) -> None:
"""Megatron 1D → 't(cp,zigzag)'."""
result: str = _megatron_plugin.infer_cp_sharded_dims(name="input_ids", ndim=1)
assert result == "t(cp,zigzag)"
def test_megatron_infer_2d(self) -> None:
"""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)"
def test_sglang_infer_1d(self) -> None:
"""SGLang 1D → 't(cp,zigzag)'."""
result: str = _sglang_plugin.infer_cp_sharded_dims(name="input_ids", ndim=1)
assert result == "t(cp,zigzag)"
def test_megatron_infer_3d_raises(self) -> None:
"""Megatron 3D raises ValueError."""
with pytest.raises(ValueError, match="cannot infer dims"):
_megatron_plugin.infer_cp_sharded_dims(name="input_ids", ndim=3)
def test_sglang_infer_2d_raises(self) -> None:
"""SGLang 2D raises ValueError."""
with pytest.raises(ValueError, match="cannot infer dims"):
_sglang_plugin.infer_cp_sharded_dims(name="input_ids", ndim=2)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))