Support overriding and post-hoc providing metadata in dump comparator (#19598)

This commit is contained in:
fzyzcjy
2026-03-01 10:35:06 +08:00
committed by GitHub
parent e41164af1c
commit e78f1283f7
5 changed files with 817 additions and 101 deletions
@@ -1,4 +1,5 @@
import sys
import textwrap
from argparse import Namespace
from pathlib import Path
@@ -1044,107 +1045,6 @@ class TestEntrypointGroupingLogical:
comp = _assert_single_comparison_passed(records)
assert comp.name == "hidden"
def test_tp_partial_reduction_unshard(self, tmp_path, capsys):
"""TP=2 with partial reduction: element-wise sum reconstructs full tensor."""
torch.manual_seed(42)
full_baseline = torch.randn(4, 8)
full_target = full_baseline + torch.randn(4, 8) * 0.001
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
baseline_path = _create_tp_partial_dumps(
baseline_dir,
full_tensor=full_baseline,
name="attn_out",
tp_size=2,
dims_str="b h(tp,partial)",
)
target_path = _create_tp_partial_dumps(
target_dir,
full_tensor=full_target,
name="attn_out",
tp_size=2,
dims_str="b h(tp,partial)",
)
args = _make_args(baseline_path, target_path, diff_threshold=0.01)
records = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records)
assert comp.name == "attn_out"
summary = records[-1]
assert isinstance(summary, SummaryRecord)
assert summary.total == 1
assert summary.passed == 1
def test_tp_partial_vs_single_rank(self, tmp_path, capsys):
"""Baseline single rank vs target TP=2 partial: unshard target then compare."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8)
target_full = full_tensor + torch.randn(4, 8) * 0.001
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
baseline_path = _create_rank_dump(
baseline_dir, rank=0, name="attn_out", tensor=full_tensor
)
target_path = _create_tp_partial_dumps(
target_dir,
full_tensor=target_full,
name="attn_out",
tp_size=2,
dims_str="b h(tp,partial)",
)
args = _make_args(baseline_path, target_path, diff_threshold=0.01)
records = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records)
assert comp.name == "attn_out"
def test_cp_concat_tp_partial_reduction(self, tmp_path, capsys):
"""CP=2 concat + TP=2 partial reduction: multi-axis unshard."""
torch.manual_seed(42)
full_baseline = torch.randn(4, 8, 16)
full_target = full_baseline + torch.randn(4, 8, 16) * 0.001
for side_dir, full_tensor in [
(tmp_path / "baseline", full_baseline),
(tmp_path / "target", full_target),
]:
side_dir.mkdir()
cp_chunks = list(full_tensor.chunk(2, dim=1))
rank = 0
for cp_rank in range(2):
for tp_rank in range(2):
_create_rank_dump(
side_dir,
rank=rank,
name="hidden",
tensor=cp_chunks[cp_rank] / 2,
dims="b s(cp) h(tp,partial)",
parallel_info={
"cp_rank": cp_rank,
"cp_size": 2,
"tp_rank": tp_rank,
"tp_size": 2,
},
)
rank += 1
args = _make_args(
tmp_path / "baseline" / _FIXED_EXP_NAME,
tmp_path / "target" / _FIXED_EXP_NAME,
diff_threshold=0.01,
)
records = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records)
assert comp.name == "hidden"
class TestEntrypointAxisAligner:
"""Test cross-framework dim reordering through the full entrypoint pipeline."""
@@ -1988,6 +1888,10 @@ def _make_args(baseline_path: Path, target_path: Path, **overrides) -> Namespace
viz_bundle_details=False,
viz_output_dir="/tmp/comparator_viz/",
visualize_per_token=None,
override_dims=[],
override_baseline_dims=[],
override_target_dims=[],
override_config=None,
)
defaults.update(overrides)
return Namespace(**defaults)
@@ -2837,5 +2741,353 @@ class TestEntrypointDpFilter:
_run_and_parse(args, capsys)
class TestEntrypointMetaOverride:
"""E2E: dump with wrong dims → --override-dims / --override-config corrects at comparison time."""
@staticmethod
def _create_single_rank_pair(
tmp_path: Path,
*,
name: str = "hidden",
baseline_dims: str | None = "x y",
target_dims: str | None = "x y",
) -> tuple[Path, Path]:
"""Create single-rank baseline+target dumps with a close tensor pair."""
torch.manual_seed(42)
tensor: torch.Tensor = torch.randn(10, 8)
target: torch.Tensor = tensor + torch.randn(10, 8) * 0.001
baseline_dir: Path = tmp_path / "baseline"
target_dir: Path = tmp_path / "target"
baseline_dir.mkdir()
target_dir.mkdir()
_create_rank_dump(
baseline_dir, rank=0, name=name, tensor=tensor, dims=baseline_dims
)
_create_rank_dump(
target_dir, rank=0, name=name, tensor=target, dims=target_dims
)
return baseline_dir / _FIXED_EXP_NAME, target_dir / _FIXED_EXP_NAME
@staticmethod
def _assert_all_passed(
records: list[AnyRecord], *, expected_count: int = 1
) -> None:
"""Assert that exactly expected_count comparisons exist and all passed."""
comparisons: list[ComparisonRecord] = _get_comparisons(records)
assert len(comparisons) == expected_count
assert all(c.diff is not None and c.diff.passed for c in comparisons)
def test_override_dims_fixes_wrong_dims(self, tmp_path: Path, capsys) -> None:
"""Tensor dumped with wrong dims='h d' is fixed by --override-dims to 't h(tp)'."""
torch.manual_seed(42)
full_tensor: torch.Tensor = torch.randn(10, 8)
tp_chunks: list[torch.Tensor] = list(full_tensor.chunk(2, dim=1))
target_full: torch.Tensor = full_tensor + torch.randn(10, 8) * 0.001
target_tp_chunks: list[torch.Tensor] = list(target_full.chunk(2, dim=1))
baseline_dir: Path = tmp_path / "baseline"
target_dir: Path = tmp_path / "target"
baseline_dir.mkdir()
target_dir.mkdir()
# Dump with WRONG dims "h d" instead of correct "t h(tp)"
for tp_rank in range(2):
_create_rank_dump(
baseline_dir,
rank=tp_rank,
name="hidden",
tensor=tp_chunks[tp_rank],
dims="h d",
parallel_info={"tp_rank": tp_rank, "tp_size": 2},
)
_create_rank_dump(
target_dir,
rank=tp_rank,
name="hidden",
tensor=target_tp_chunks[tp_rank],
dims="h d",
parallel_info={"tp_rank": tp_rank, "tp_size": 2},
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
grouping="logical",
override_dims=["hidden:t h(tp)"],
)
self._assert_all_passed(_run_and_parse(args, capsys))
@pytest.mark.parametrize(
"baseline_dims, target_dims, override_kwarg",
[
("x y", "t h", {"override_baseline_dims": ["hidden:t h"]}),
("t h", "x y", {"override_target_dims": ["hidden:t h"]}),
("x y", "x y", {"override_dims": ["hidden:t h"]}),
],
ids=["baseline_only", "target_only", "both_via_override_dims"],
)
def test_single_side_override(
self,
tmp_path: Path,
capsys,
baseline_dims: str,
target_dims: str,
override_kwarg: dict,
) -> None:
"""Per-side override fixes the wrong dims on one or both sides."""
baseline_path, target_path = self._create_single_rank_pair(
tmp_path,
baseline_dims=baseline_dims,
target_dims=target_dims,
)
args = _make_args(baseline_path, target_path, grouping="raw", **override_kwarg)
self._assert_all_passed(_run_and_parse(args, capsys))
def test_override_config_yaml(self, tmp_path: Path, capsys) -> None:
"""--override-config YAML overrides dims."""
baseline_path, target_path = self._create_single_rank_pair(tmp_path)
yaml_path: Path = tmp_path / "override.yaml"
yaml_path.write_text(textwrap.dedent("""\
overrides:
- match: "hidden"
dims: "t h"
"""))
args = _make_args(
baseline_path,
target_path,
grouping="raw",
override_config=str(yaml_path),
)
self._assert_all_passed(_run_and_parse(args, capsys))
def test_no_match_uses_original_dims(self, tmp_path: Path, capsys) -> None:
"""When override regex doesn't match, original dims from dump are used."""
baseline_path, target_path = self._create_single_rank_pair(
tmp_path,
baseline_dims="t h",
target_dims="t h",
)
args = _make_args(
baseline_path,
target_path,
grouping="raw",
override_dims=["no_match_pattern:b s d"],
)
self._assert_all_passed(_run_and_parse(args, capsys))
def test_selective_match_multi_tensor(self, tmp_path: Path, capsys) -> None:
"""Override matches only 'logits'; 'hidden' uses original dims."""
torch.manual_seed(42)
baseline_dir: Path = tmp_path / "baseline"
target_dir: Path = tmp_path / "target"
baseline_dir.mkdir()
target_dir.mkdir()
hidden_b: torch.Tensor = torch.randn(10, 8)
hidden_t: torch.Tensor = hidden_b + torch.randn(10, 8) * 0.001
logits_b: torch.Tensor = torch.randn(10, 4)
logits_t: torch.Tensor = logits_b + torch.randn(10, 4) * 0.001
for name, b_tensor, t_tensor, dims in [
("hidden", hidden_b, hidden_t, "t h"),
("logits", logits_b, logits_t, "x y"),
]:
_create_rank_dump(
baseline_dir, rank=0, name=name, tensor=b_tensor, dims=dims
)
_create_rank_dump(target_dir, rank=0, name=name, tensor=t_tensor, dims=dims)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
grouping="raw",
override_dims=["logits:t v"],
)
self._assert_all_passed(_run_and_parse(args, capsys), expected_count=2)
def test_multiple_cli_override_dims(self, tmp_path: Path, capsys) -> None:
"""Multiple --override-dims for different tensors."""
torch.manual_seed(42)
baseline_dir: Path = tmp_path / "baseline"
target_dir: Path = tmp_path / "target"
baseline_dir.mkdir()
target_dir.mkdir()
hidden_b: torch.Tensor = torch.randn(10, 8)
hidden_t: torch.Tensor = hidden_b + torch.randn(10, 8) * 0.001
logits_b: torch.Tensor = torch.randn(10, 4)
logits_t: torch.Tensor = logits_b + torch.randn(10, 4) * 0.001
for name, b_tensor, t_tensor in [
("hidden", hidden_b, hidden_t),
("logits", logits_b, logits_t),
]:
_create_rank_dump(
baseline_dir, rank=0, name=name, tensor=b_tensor, dims="x y"
)
_create_rank_dump(
target_dir, rank=0, name=name, tensor=t_tensor, dims="x y"
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
grouping="raw",
override_dims=["hidden:t h", "logits:t v"],
)
self._assert_all_passed(_run_and_parse(args, capsys), expected_count=2)
def test_per_side_dims_different_parallelism(self, tmp_path: Path, capsys) -> None:
"""baseline TP-sharded, target EP-sharded — per-side override fixes both."""
torch.manual_seed(42)
full_tensor: torch.Tensor = torch.randn(10, 8)
target_full: torch.Tensor = full_tensor + torch.randn(10, 8) * 0.001
baseline_dir: Path = tmp_path / "baseline"
target_dir: Path = tmp_path / "target"
baseline_dir.mkdir()
target_dir.mkdir()
b_chunks: list[torch.Tensor] = list(full_tensor.chunk(2, dim=1))
for tp_rank in range(2):
_create_rank_dump(
baseline_dir,
rank=tp_rank,
name="hidden",
tensor=b_chunks[tp_rank],
dims="x y",
parallel_info={"tp_rank": tp_rank, "tp_size": 2},
)
t_chunks: list[torch.Tensor] = list(target_full.chunk(2, dim=1))
for ep_rank in range(2):
_create_rank_dump(
target_dir,
rank=ep_rank,
name="hidden",
tensor=t_chunks[ep_rank],
dims="x y",
parallel_info={"ep_rank": ep_rank, "ep_size": 2},
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
grouping="logical",
override_baseline_dims=["hidden:t h(tp)"],
override_target_dims=["hidden:t h(ep)"],
)
self._assert_all_passed(_run_and_parse(args, capsys))
def test_yaml_first_match_wins_e2e(self, tmp_path: Path, capsys) -> None:
"""YAML with two matching rules: first rule wins in real pipeline."""
baseline_path, target_path = self._create_single_rank_pair(tmp_path)
yaml_path: Path = tmp_path / "override.yaml"
yaml_path.write_text(textwrap.dedent("""\
overrides:
- match: "hidden"
dims: "t h"
- match: "hidden"
dims: "a b"
"""))
args = _make_args(
baseline_path,
target_path,
grouping="raw",
override_config=str(yaml_path),
)
self._assert_all_passed(_run_and_parse(args, capsys))
def test_cli_overrides_yaml_e2e(self, tmp_path: Path, capsys) -> None:
"""CLI --override-dims wins over YAML rule for the same tensor."""
baseline_path, target_path = self._create_single_rank_pair(tmp_path)
yaml_path: Path = tmp_path / "override.yaml"
yaml_path.write_text(textwrap.dedent("""\
overrides:
- match: "hidden"
dims: "a b"
"""))
args = _make_args(
baseline_path,
target_path,
grouping="raw",
override_dims=["hidden:t h"],
override_config=str(yaml_path),
)
self._assert_all_passed(_run_and_parse(args, capsys))
def test_override_injects_dims_when_absent(self, tmp_path: Path, capsys) -> None:
"""Override injects dims into meta even when dump had no dims annotation."""
baseline_path, target_path = self._create_single_rank_pair(
tmp_path,
baseline_dims=None,
target_dims=None,
)
args = _make_args(
baseline_path,
target_path,
grouping="raw",
override_dims=["hidden:t h"],
)
self._assert_all_passed(_run_and_parse(args, capsys))
def test_non_tensor_unaffected_by_override(self, tmp_path: Path, capsys) -> None:
"""Non-tensor values pass through without error even with active override."""
torch.manual_seed(42)
tensor: torch.Tensor = torch.randn(4, 4)
baseline_dir: Path = tmp_path / "baseline"
target_dir: Path = tmp_path / "target"
baseline_dir.mkdir()
target_dir.mkdir()
for side_dir in [baseline_dir, target_dir]:
_create_non_tensor_rank_dump(
side_dir,
rank=0,
name="sm_scale",
value=0.125,
extra_tensor_dumps=[("hidden", tensor)],
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
grouping="raw",
override_dims=["hidden:x y"],
)
records = _run_and_parse(args, capsys)
non_tensors: list[NonTensorRecord] = [
r for r in records if isinstance(r, NonTensorRecord)
]
assert len(non_tensors) == 1
assert non_tensors[0].name == "sm_scale"
assert non_tensors[0].values_equal
comparisons: list[ComparisonRecord] = _get_comparisons(records)
assert len(comparisons) == 1
assert comparisons[0].name == "hidden"
summary: SummaryRecord = [r for r in records if isinstance(r, SummaryRecord)][0]
assert summary.failed == 0
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))