Support handling arbitrary objects in dump comparator (#19558)
This commit is contained in:
@@ -21,6 +21,7 @@ from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
|
||||
from sglang.srt.debug_utils.comparator.dims import apply_dim_names, parse_dim_names
|
||||
from sglang.srt.debug_utils.comparator.output_types import (
|
||||
ComparisonRecord,
|
||||
NonTensorRecord,
|
||||
SkipRecord,
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.tensor_comparator.comparator import (
|
||||
@@ -28,7 +29,7 @@ from sglang.srt.debug_utils.comparator.tensor_comparator.comparator import (
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.utils import Pair
|
||||
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
|
||||
from sglang.srt.debug_utils.dump_loader import ValueWithMeta
|
||||
from sglang.srt.debug_utils.dump_loader import LOAD_FAILED, ValueWithMeta
|
||||
|
||||
_FAILED_SIDE_MAP: dict[str, str] = {"x": "baseline", "y": "target"}
|
||||
|
||||
@@ -44,9 +45,9 @@ def compare_bundle_pair(
|
||||
thd_seq_lens_by_step_pair: Pair[Optional[dict[int, list[int]]]] = Pair(
|
||||
x=None, y=None
|
||||
),
|
||||
) -> Union[ComparisonRecord, SkipRecord]:
|
||||
) -> Union[ComparisonRecord, SkipRecord, NonTensorRecord]:
|
||||
with warning_sink.context() as collected_warnings:
|
||||
result = _compare_bundle_pair_raw(
|
||||
result = _compare_bundle_pair_inner(
|
||||
name=name,
|
||||
filenames_pair=filenames_pair,
|
||||
baseline_path=baseline_path,
|
||||
@@ -59,7 +60,7 @@ def compare_bundle_pair(
|
||||
return result.model_copy(update={"warnings": collected_warnings})
|
||||
|
||||
|
||||
def _compare_bundle_pair_raw(
|
||||
def _compare_bundle_pair_inner(
|
||||
*,
|
||||
name: str,
|
||||
filenames_pair: Pair[list[str]],
|
||||
@@ -70,18 +71,49 @@ def _compare_bundle_pair_raw(
|
||||
thd_seq_lens_by_step_pair: Pair[Optional[dict[int, list[int]]]] = Pair(
|
||||
x=None, y=None
|
||||
),
|
||||
) -> Union[ComparisonRecord, SkipRecord]:
|
||||
# 1. Load (tensor + meta, ungrouped)
|
||||
valid_pair: Pair[list[ValueWithMeta]] = Pair(
|
||||
x=_load_valid_tensors(filenames=filenames_pair.x, base_path=baseline_path),
|
||||
y=_load_valid_tensors(filenames=filenames_pair.y, base_path=target_path),
|
||||
) -> Union[ComparisonRecord, SkipRecord, NonTensorRecord]:
|
||||
# 1. Load all successfully loaded values
|
||||
all_pair: Pair[list[ValueWithMeta]] = Pair(
|
||||
x=_load_all_values(filenames=filenames_pair.x, base_path=baseline_path),
|
||||
y=_load_all_values(filenames=filenames_pair.y, base_path=target_path),
|
||||
)
|
||||
|
||||
if not all_pair.x or not all_pair.y:
|
||||
reason = "baseline_load_failed" if not all_pair.x else "target_load_failed"
|
||||
return SkipRecord(name=name, reason=reason)
|
||||
|
||||
# 2. Check if any side has non-tensor values → non-tensor display path
|
||||
has_non_tensor: bool = any(
|
||||
not isinstance(it.value, torch.Tensor) for it in [*all_pair.x, *all_pair.y]
|
||||
)
|
||||
if has_non_tensor:
|
||||
return _compare_bundle_pair_non_tensor_type(name=name, value_pair=all_pair)
|
||||
|
||||
# 3. All values are tensors → tensor comparison path
|
||||
return _compare_bundle_pair_tensor_type(
|
||||
name=name,
|
||||
valid_pair=all_pair,
|
||||
token_aligner_plan=token_aligner_plan,
|
||||
diff_threshold=diff_threshold,
|
||||
thd_seq_lens_by_step_pair=thd_seq_lens_by_step_pair,
|
||||
)
|
||||
|
||||
|
||||
def _compare_bundle_pair_tensor_type(
|
||||
*,
|
||||
name: str,
|
||||
valid_pair: Pair[list[ValueWithMeta]],
|
||||
token_aligner_plan: Optional[TokenAlignerPlan],
|
||||
diff_threshold: float,
|
||||
thd_seq_lens_by_step_pair: Pair[Optional[dict[int, list[int]]]] = Pair(
|
||||
x=None, y=None
|
||||
),
|
||||
) -> Union[ComparisonRecord, SkipRecord]:
|
||||
if not valid_pair.x or not valid_pair.y:
|
||||
reason = "baseline_load_failed" if not valid_pair.x else "target_load_failed"
|
||||
return SkipRecord(name=name, reason=reason)
|
||||
|
||||
# 2. Plan (meta only, no tensor)
|
||||
# Plan (meta only, no tensor)
|
||||
metas_pair: Pair[list[dict[str, Any]]] = valid_pair.map(
|
||||
lambda items: [it.meta for it in items]
|
||||
)
|
||||
@@ -91,7 +123,7 @@ def _compare_bundle_pair_raw(
|
||||
thd_seq_lens_by_step_pair=thd_seq_lens_by_step_pair,
|
||||
)
|
||||
|
||||
# 3. Apply dim names to tensors, then execute
|
||||
# Apply dim names to tensors, then execute
|
||||
tensors_pair: Pair[list[torch.Tensor]] = Pair(
|
||||
x=_apply_dim_names_from_meta(
|
||||
tensors=[it.value for it in valid_pair.x],
|
||||
@@ -109,10 +141,10 @@ def _compare_bundle_pair_raw(
|
||||
if aligner_result.tensors is None:
|
||||
assert aligner_result.failed_side_xy is not None
|
||||
side_name: str = _FAILED_SIDE_MAP[aligner_result.failed_side_xy]
|
||||
reason = f"{side_name}_load_failed"
|
||||
reason: str = f"{side_name}_load_failed"
|
||||
return SkipRecord(name=name, reason=reason)
|
||||
|
||||
# 4. Compare
|
||||
# Compare
|
||||
info = compare_tensor_pair(
|
||||
x_baseline=aligner_result.tensors.x.rename(None),
|
||||
x_target=aligner_result.tensors.y.rename(None),
|
||||
@@ -122,20 +154,27 @@ def _compare_bundle_pair_raw(
|
||||
return ComparisonRecord(**info.model_dump(), aligner_plan=plan)
|
||||
|
||||
|
||||
def _apply_dim_names_from_meta(
|
||||
def _compare_bundle_pair_non_tensor_type(
|
||||
*,
|
||||
tensors: list[torch.Tensor],
|
||||
metas: list[dict[str, Any]],
|
||||
) -> list[torch.Tensor]:
|
||||
if not metas:
|
||||
return tensors
|
||||
name: str,
|
||||
value_pair: Pair[list[ValueWithMeta]],
|
||||
) -> NonTensorRecord:
|
||||
baseline_value: Any = value_pair.x[0].value
|
||||
target_value: Any = value_pair.y[0].value
|
||||
|
||||
dims_str: Optional[str] = metas[0].get("dims")
|
||||
if dims_str is None:
|
||||
return tensors
|
||||
try:
|
||||
values_equal: bool = bool(baseline_value == target_value)
|
||||
except Exception:
|
||||
values_equal = False
|
||||
|
||||
dim_names: list[str] = parse_dim_names(dims_str)
|
||||
return [apply_dim_names(t, dim_names) for t in tensors]
|
||||
return NonTensorRecord(
|
||||
name=name,
|
||||
baseline_value=repr(baseline_value),
|
||||
target_value=repr(target_value),
|
||||
baseline_type=type(baseline_value).__name__,
|
||||
target_type=type(target_value).__name__,
|
||||
values_equal=values_equal,
|
||||
)
|
||||
|
||||
|
||||
def _apply_dim_names_from_meta(
|
||||
@@ -154,9 +193,9 @@ def _apply_dim_names_from_meta(
|
||||
return [apply_dim_names(t, dim_names) for t in tensors]
|
||||
|
||||
|
||||
def _load_valid_tensors(filenames: list[str], base_path: Path) -> list[ValueWithMeta]:
|
||||
def _load_all_values(filenames: list[str], base_path: Path) -> list[ValueWithMeta]:
|
||||
return [
|
||||
x
|
||||
item
|
||||
for f in filenames
|
||||
if isinstance((x := ValueWithMeta.load(base_path / f)).value, torch.Tensor)
|
||||
if (item := ValueWithMeta.load(base_path / f)).value is not LOAD_FAILED
|
||||
]
|
||||
|
||||
@@ -12,7 +12,7 @@ from sglang.srt.debug_utils.comparator.output_types import (
|
||||
RankInfoRecord,
|
||||
print_record,
|
||||
)
|
||||
from sglang.srt.debug_utils.dump_loader import ValueWithMeta
|
||||
from sglang.srt.debug_utils.dump_loader import LOAD_FAILED, ValueWithMeta
|
||||
|
||||
_PARALLEL_INFO_KEYS: list[str] = ["sglang_parallel_info", "megatron_parallel_info"]
|
||||
|
||||
@@ -96,7 +96,7 @@ def _collect_input_ids_and_positions(
|
||||
for row in filtered.to_dicts():
|
||||
key: tuple[int, int] = (row["step"], row["rank"])
|
||||
item: ValueWithMeta = ValueWithMeta.load(dump_dir / row["filename"])
|
||||
if item.value is not None:
|
||||
if item.value is not LOAD_FAILED:
|
||||
data_by_step_rank[key][row["name"]] = item.value
|
||||
|
||||
table_rows: list[dict[str, Any]] = []
|
||||
|
||||
@@ -25,6 +25,7 @@ from sglang.srt.debug_utils.comparator.display import emit_display_records
|
||||
from sglang.srt.debug_utils.comparator.output_types import (
|
||||
ComparisonRecord,
|
||||
ConfigRecord,
|
||||
NonTensorRecord,
|
||||
SkipRecord,
|
||||
SummaryRecord,
|
||||
print_record,
|
||||
@@ -137,7 +138,7 @@ def _compare_bundle_pairs(
|
||||
token_aligner_plan: Optional[TokenAlignerPlan],
|
||||
diff_threshold: float,
|
||||
thd_seq_lens_by_step_pair: Pair[Optional[dict[int, list[int]]]],
|
||||
) -> Iterator[Union[ComparisonRecord, SkipRecord]]:
|
||||
) -> Iterator[Union[ComparisonRecord, SkipRecord, NonTensorRecord]]:
|
||||
for bundle_info_pair in bundle_info_pairs:
|
||||
if not bundle_info_pair.y:
|
||||
continue
|
||||
@@ -159,7 +160,7 @@ def _compare_bundle_pairs(
|
||||
|
||||
def _consume_comparison_records(
|
||||
*,
|
||||
comparison_records: Iterator[Union[ComparisonRecord, SkipRecord]],
|
||||
comparison_records: Iterator[Union[ComparisonRecord, SkipRecord, NonTensorRecord]],
|
||||
output_format: str,
|
||||
) -> None:
|
||||
counts: dict[str, int] = {"passed": 0, "failed": 0, "skipped": 0}
|
||||
|
||||
@@ -140,6 +140,31 @@ class ComparisonRecord(TensorComparisonInfo, _OutputRecord):
|
||||
return body
|
||||
|
||||
|
||||
class NonTensorRecord(_OutputRecord):
|
||||
type: Literal["non_tensor"] = "non_tensor"
|
||||
name: str
|
||||
baseline_value: str
|
||||
target_value: str
|
||||
baseline_type: str
|
||||
target_type: str
|
||||
values_equal: bool
|
||||
|
||||
@property
|
||||
def category(self) -> str:
|
||||
if self.warnings:
|
||||
return "failed"
|
||||
return "passed" if self.values_equal else "failed"
|
||||
|
||||
def _format_body(self) -> str:
|
||||
if self.values_equal:
|
||||
return f"NonTensor: {self.name} = {self.baseline_value} ({self.baseline_type}) [equal]"
|
||||
return (
|
||||
f"NonTensor: {self.name}\n"
|
||||
f" baseline = {self.baseline_value} ({self.baseline_type})\n"
|
||||
f" target = {self.target_value} ({self.target_type})"
|
||||
)
|
||||
|
||||
|
||||
class SummaryRecord(_OutputRecord):
|
||||
type: Literal["summary"] = "summary"
|
||||
total: int
|
||||
@@ -207,6 +232,7 @@ AnyRecord = Annotated[
|
||||
InputIdsRecord,
|
||||
SkipRecord,
|
||||
ComparisonRecord,
|
||||
NonTensorRecord,
|
||||
SummaryRecord,
|
||||
WarningRecord,
|
||||
],
|
||||
|
||||
@@ -9,6 +9,8 @@ import torch
|
||||
|
||||
_TYPED_FIELDS: list[tuple[str, type]] = [("rank", int)]
|
||||
|
||||
LOAD_FAILED: object = object()
|
||||
|
||||
|
||||
def parse_meta_from_filename(path: Path) -> Dict[str, Any]:
|
||||
stem = Path(path).stem
|
||||
@@ -38,7 +40,7 @@ class ValueWithMeta:
|
||||
except Exception as e:
|
||||
print(f"Skip load {path} since error {e}")
|
||||
return ValueWithMeta(
|
||||
value=None, meta={**meta_from_filename, "filename": path.name}
|
||||
value=LOAD_FAILED, meta={**meta_from_filename, "filename": path.name}
|
||||
)
|
||||
|
||||
value, meta_from_embedded = _unwrap_dict_format(raw)
|
||||
|
||||
@@ -12,6 +12,7 @@ from sglang.srt.debug_utils.comparator.output_types import (
|
||||
ComparisonRecord,
|
||||
ConfigRecord,
|
||||
GeneralWarning,
|
||||
NonTensorRecord,
|
||||
SkipRecord,
|
||||
SummaryRecord,
|
||||
WarningRecord,
|
||||
@@ -1303,6 +1304,147 @@ class TestEntrypointAlignment:
|
||||
assert summary.passed == 2
|
||||
|
||||
|
||||
class TestEntrypointNonTensorValues:
|
||||
"""Test non-tensor value comparison through the full entrypoint pipeline."""
|
||||
|
||||
def test_non_tensor_float_same_value(self, tmp_path: Path, capsys) -> None:
|
||||
"""Two sides dump the same float → NonTensorRecord with values_equal=True, category=passed."""
|
||||
baseline_path, target_path = _create_non_tensor_dumps(
|
||||
tmp_path, name="sm_scale", baseline_value=0.125, target_value=0.125
|
||||
)
|
||||
args = _make_args(baseline_path, target_path, grouping="raw")
|
||||
records = _run_and_parse(args, capsys)
|
||||
|
||||
non_tensors = _get_non_tensors(records)
|
||||
assert len(non_tensors) == 1
|
||||
assert non_tensors[0].name == "sm_scale"
|
||||
assert non_tensors[0].values_equal is True
|
||||
assert non_tensors[0].category == "passed"
|
||||
|
||||
summary = records[-1]
|
||||
assert isinstance(summary, SummaryRecord)
|
||||
assert summary.passed == 1
|
||||
assert summary.failed == 0
|
||||
|
||||
def test_non_tensor_float_different_value(self, tmp_path: Path, capsys) -> None:
|
||||
"""Two sides dump different floats → NonTensorRecord with values_equal=False, category=failed."""
|
||||
baseline_path, target_path = _create_non_tensor_dumps(
|
||||
tmp_path, name="sm_scale", baseline_value=0.125, target_value=0.25
|
||||
)
|
||||
args = _make_args(baseline_path, target_path, grouping="raw")
|
||||
records = _run_and_parse(args, capsys)
|
||||
|
||||
non_tensors = _get_non_tensors(records)
|
||||
assert len(non_tensors) == 1
|
||||
assert non_tensors[0].values_equal is False
|
||||
assert non_tensors[0].category == "failed"
|
||||
|
||||
summary = records[-1]
|
||||
assert isinstance(summary, SummaryRecord)
|
||||
assert summary.failed == 1
|
||||
|
||||
def test_non_tensor_string_value(self, tmp_path: Path, capsys) -> None:
|
||||
"""String non-tensor values are compared and displayed correctly."""
|
||||
baseline_path, target_path = _create_non_tensor_dumps(
|
||||
tmp_path,
|
||||
name="attn_backend",
|
||||
baseline_value="flash_attn",
|
||||
target_value="flash_attn",
|
||||
)
|
||||
args = _make_args(baseline_path, target_path, grouping="raw")
|
||||
records = _run_and_parse(args, capsys)
|
||||
|
||||
non_tensors = _get_non_tensors(records)
|
||||
assert len(non_tensors) == 1
|
||||
assert non_tensors[0].values_equal is True
|
||||
assert non_tensors[0].baseline_type == "str"
|
||||
assert non_tensors[0].target_type == "str"
|
||||
|
||||
def test_non_tensor_mixed_with_tensor(self, tmp_path: Path, capsys) -> None:
|
||||
"""Tensors and non_tensors in the same dump are each handled correctly."""
|
||||
torch.manual_seed(42)
|
||||
tensor = torch.randn(4, 4)
|
||||
|
||||
baseline_dir = tmp_path / "baseline"
|
||||
target_dir = tmp_path / "target"
|
||||
|
||||
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",
|
||||
)
|
||||
records = _run_and_parse(args, capsys)
|
||||
|
||||
comparisons = _get_comparisons(records)
|
||||
non_tensors = _get_non_tensors(records)
|
||||
assert len(comparisons) == 1
|
||||
assert comparisons[0].name == "hidden"
|
||||
assert len(non_tensors) == 1
|
||||
assert non_tensors[0].name == "sm_scale"
|
||||
assert non_tensors[0].values_equal is True
|
||||
|
||||
summary = records[-1]
|
||||
assert isinstance(summary, SummaryRecord)
|
||||
assert summary.passed == 2
|
||||
|
||||
def test_non_tensor_complex_object(self, tmp_path: Path, capsys) -> None:
|
||||
"""Complex objects (e.g. dict containing a tensor) are displayed via repr, not skipped."""
|
||||
value = {"a": 1, "b": "hello", "c": torch.tensor([1.0, 2.0])}
|
||||
baseline_path, target_path = _create_non_tensor_dumps(
|
||||
tmp_path, name="debug_info", baseline_value=value, target_value=value
|
||||
)
|
||||
args = _make_args(baseline_path, target_path, grouping="raw")
|
||||
records = _run_and_parse(args, capsys)
|
||||
|
||||
non_tensors = _get_non_tensors(records)
|
||||
assert len(non_tensors) == 1
|
||||
assert non_tensors[0].name == "debug_info"
|
||||
assert non_tensors[0].baseline_type == "dict"
|
||||
assert non_tensors[0].target_type == "dict"
|
||||
|
||||
def test_non_tensor_none_value(self, tmp_path: Path, capsys) -> None:
|
||||
"""Dumping None is displayed as NonTensorRecord, not skipped as load failure."""
|
||||
baseline_path, target_path = _create_non_tensor_dumps(
|
||||
tmp_path, name="optional_param", baseline_value=None, target_value=None
|
||||
)
|
||||
args = _make_args(baseline_path, target_path, grouping="raw")
|
||||
records = _run_and_parse(args, capsys)
|
||||
|
||||
non_tensors = _get_non_tensors(records)
|
||||
assert len(non_tensors) == 1
|
||||
assert non_tensors[0].name == "optional_param"
|
||||
assert non_tensors[0].values_equal is True
|
||||
assert non_tensors[0].baseline_value == "None"
|
||||
assert non_tensors[0].baseline_type == "NoneType"
|
||||
assert non_tensors[0].category == "passed"
|
||||
|
||||
def test_non_tensor_json_roundtrip(self, tmp_path: Path, capsys) -> None:
|
||||
"""NonTensorRecord JSON output can be parsed back correctly."""
|
||||
baseline_path, target_path = _create_non_tensor_dumps(
|
||||
tmp_path, name="sm_scale", baseline_value=0.125, target_value=0.125
|
||||
)
|
||||
args = _make_args(baseline_path, target_path, grouping="raw")
|
||||
records = _run_and_parse(args, capsys)
|
||||
|
||||
non_tensors = _get_non_tensors(records)
|
||||
assert len(non_tensors) == 1
|
||||
|
||||
json_str: str = non_tensors[0].model_dump_json()
|
||||
roundtripped = parse_record_json(json_str)
|
||||
assert isinstance(roundtripped, NonTensorRecord)
|
||||
assert roundtripped.name == "sm_scale"
|
||||
assert roundtripped.values_equal is True
|
||||
|
||||
|
||||
# --------------------------- Assertion helpers -------------------
|
||||
|
||||
|
||||
@@ -1310,6 +1452,10 @@ def _get_comparisons(records: list[AnyRecord]) -> list[ComparisonRecord]:
|
||||
return [r for r in records if isinstance(r, ComparisonRecord)]
|
||||
|
||||
|
||||
def _get_non_tensors(records: list[AnyRecord]) -> list[NonTensorRecord]:
|
||||
return [r for r in records if isinstance(r, NonTensorRecord)]
|
||||
|
||||
|
||||
def _assert_single_comparison_passed(records: list[AnyRecord]) -> ComparisonRecord:
|
||||
comparisons = _get_comparisons(records)
|
||||
assert len(comparisons) == 1
|
||||
@@ -1366,6 +1512,56 @@ def _create_dumps(
|
||||
return exp_paths[0], exp_paths[1]
|
||||
|
||||
|
||||
def _create_non_tensor_rank_dump(
|
||||
directory: Path,
|
||||
*,
|
||||
rank: int,
|
||||
name: str,
|
||||
value: object,
|
||||
extra_tensor_dumps: list[tuple[str, torch.Tensor]] | None = None,
|
||||
) -> Path:
|
||||
with pytest.MonkeyPatch.context() as mp:
|
||||
mp.setattr(_dumper_module, "_get_rank", lambda: rank)
|
||||
|
||||
dumper = _Dumper(
|
||||
config=DumperConfig(
|
||||
enable=True,
|
||||
dir=str(directory),
|
||||
exp_name=_FIXED_EXP_NAME,
|
||||
enable_http_server=False,
|
||||
)
|
||||
)
|
||||
dumper.__dict__["_static_meta"] = {"world_rank": rank, "world_size": 1}
|
||||
|
||||
dumper.dump(name, value)
|
||||
for extra_name, extra_tensor in extra_tensor_dumps or []:
|
||||
dumper.dump(extra_name, extra_tensor)
|
||||
dumper.step()
|
||||
|
||||
return directory / _FIXED_EXP_NAME
|
||||
|
||||
|
||||
def _create_non_tensor_dumps(
|
||||
tmp_path: Path,
|
||||
*,
|
||||
name: str,
|
||||
baseline_value: object,
|
||||
target_value: object,
|
||||
) -> tuple[Path, Path]:
|
||||
baseline_dir = tmp_path / "baseline"
|
||||
target_dir = tmp_path / "target"
|
||||
baseline_dir.mkdir()
|
||||
target_dir.mkdir()
|
||||
|
||||
baseline_path = _create_non_tensor_rank_dump(
|
||||
baseline_dir, rank=0, name=name, value=baseline_value
|
||||
)
|
||||
target_path = _create_non_tensor_rank_dump(
|
||||
target_dir, rank=0, name=name, value=target_value
|
||||
)
|
||||
return baseline_path, target_path
|
||||
|
||||
|
||||
def _make_args(baseline_path: Path, target_path: Path, **overrides) -> Namespace:
|
||||
defaults = dict(
|
||||
baseline_path=str(baseline_path),
|
||||
|
||||
@@ -24,6 +24,7 @@ from sglang.srt.debug_utils.comparator.dims import ParallelAxis, TokenLayout
|
||||
from sglang.srt.debug_utils.comparator.output_types import (
|
||||
ComparisonRecord,
|
||||
GeneralWarning,
|
||||
NonTensorRecord,
|
||||
SkipRecord,
|
||||
SummaryRecord,
|
||||
parse_record_json,
|
||||
@@ -250,6 +251,157 @@ class TestOutputRecordCategories:
|
||||
)
|
||||
assert record.category == "passed"
|
||||
|
||||
def test_non_tensor_record_equal_is_passed(self) -> None:
|
||||
record = NonTensorRecord(
|
||||
name="sm_scale",
|
||||
baseline_value="0.125",
|
||||
target_value="0.125",
|
||||
baseline_type="float",
|
||||
target_type="float",
|
||||
values_equal=True,
|
||||
)
|
||||
assert record.category == "passed"
|
||||
|
||||
def test_non_tensor_record_different_is_failed(self) -> None:
|
||||
record = NonTensorRecord(
|
||||
name="sm_scale",
|
||||
baseline_value="0.125",
|
||||
target_value="0.25",
|
||||
baseline_type="float",
|
||||
target_type="float",
|
||||
values_equal=False,
|
||||
)
|
||||
assert record.category == "failed"
|
||||
|
||||
def test_non_tensor_record_with_warnings_is_failed(self) -> None:
|
||||
record = NonTensorRecord(
|
||||
name="sm_scale",
|
||||
baseline_value="0.125",
|
||||
target_value="0.125",
|
||||
baseline_type="float",
|
||||
target_type="float",
|
||||
values_equal=True,
|
||||
warnings=[GeneralWarning(category="c", message="m")],
|
||||
)
|
||||
assert record.category == "failed"
|
||||
|
||||
def test_non_tensor_record_json_roundtrip(self) -> None:
|
||||
record = NonTensorRecord(
|
||||
name="sm_scale",
|
||||
baseline_value="0.125",
|
||||
target_value="0.25",
|
||||
baseline_type="float",
|
||||
target_type="float",
|
||||
values_equal=False,
|
||||
)
|
||||
json_str: str = record.model_dump_json()
|
||||
roundtripped = parse_record_json(json_str)
|
||||
assert isinstance(roundtripped, NonTensorRecord)
|
||||
assert roundtripped.name == "sm_scale"
|
||||
assert roundtripped.values_equal is False
|
||||
assert roundtripped.baseline_value == "0.125"
|
||||
assert roundtripped.target_value == "0.25"
|
||||
|
||||
def test_non_tensor_record_text_format_equal(self) -> None:
|
||||
record = NonTensorRecord(
|
||||
name="sm_scale",
|
||||
baseline_value="0.125",
|
||||
target_value="0.125",
|
||||
baseline_type="float",
|
||||
target_type="float",
|
||||
values_equal=True,
|
||||
)
|
||||
text: str = record.to_text()
|
||||
assert "sm_scale" in text
|
||||
assert "[equal]" in text
|
||||
|
||||
def test_non_tensor_record_text_format_different(self) -> None:
|
||||
record = NonTensorRecord(
|
||||
name="sm_scale",
|
||||
baseline_value="0.125",
|
||||
target_value="0.25",
|
||||
baseline_type="float",
|
||||
target_type="float",
|
||||
values_equal=False,
|
||||
)
|
||||
text: str = record.to_text()
|
||||
assert "baseline" in text
|
||||
assert "target" in text
|
||||
|
||||
|
||||
def _make_aligner_plan() -> AlignerPlan:
|
||||
unsharder = UnsharderPlan(
|
||||
axis=ParallelAxis.TP,
|
||||
params=ConcatParams(dim_name="h"),
|
||||
groups=[[0, 1]],
|
||||
)
|
||||
return AlignerPlan(
|
||||
per_step_plans=Pair(
|
||||
x=[
|
||||
AlignerPerStepPlan(
|
||||
step=0, input_object_indices=[0, 1], sub_plans=[unsharder]
|
||||
)
|
||||
],
|
||||
y=[
|
||||
AlignerPerStepPlan(
|
||||
step=0, input_object_indices=[0, 1], sub_plans=[unsharder]
|
||||
)
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class TestAlignerPlanInComparisonRecord:
|
||||
def test_comparison_record_with_aligner_plan(self) -> None:
|
||||
plan: AlignerPlan = _make_aligner_plan()
|
||||
record: ComparisonRecord = _make_comparison_record(
|
||||
diff=_make_diff_info(passed=True),
|
||||
)
|
||||
record_with_plan = record.model_copy(update={"aligner_plan": plan})
|
||||
assert record_with_plan.aligner_plan is not None
|
||||
assert record_with_plan.aligner_plan.per_step_plans.x[0].step == 0
|
||||
|
||||
def test_aligner_plan_json_roundtrip(self) -> None:
|
||||
plan: AlignerPlan = _make_aligner_plan()
|
||||
record: ComparisonRecord = _make_comparison_record(
|
||||
diff=_make_diff_info(passed=True),
|
||||
)
|
||||
record_with_plan = record.model_copy(update={"aligner_plan": plan})
|
||||
|
||||
json_str: str = record_with_plan.model_dump_json()
|
||||
parsed = json.loads(json_str)
|
||||
assert "aligner_plan" in parsed
|
||||
assert (
|
||||
parsed["aligner_plan"]["per_step_plans"]["x"][0]["sub_plans"][0]["type"]
|
||||
== "unsharder"
|
||||
)
|
||||
|
||||
roundtripped: ComparisonRecord = parse_record_json(json_str)
|
||||
assert roundtripped.aligner_plan is not None
|
||||
assert (
|
||||
roundtripped.aligner_plan.per_step_plans.x[0].sub_plans[0].type
|
||||
== "unsharder"
|
||||
)
|
||||
|
||||
def test_comparison_record_without_aligner_plan(self) -> None:
|
||||
record: ComparisonRecord = _make_comparison_record(
|
||||
diff=_make_diff_info(passed=True),
|
||||
)
|
||||
json_str: str = record.model_dump_json()
|
||||
roundtripped: ComparisonRecord = parse_record_json(json_str)
|
||||
assert roundtripped.aligner_plan is None
|
||||
|
||||
def test_aligner_plan_text_format(self) -> None:
|
||||
plan: AlignerPlan = _make_aligner_plan()
|
||||
record: ComparisonRecord = _make_comparison_record(
|
||||
diff=_make_diff_info(passed=True),
|
||||
)
|
||||
record_with_plan = record.model_copy(update={"aligner_plan": plan})
|
||||
|
||||
text: str = record_with_plan.to_text()
|
||||
assert "Aligner Plan:" in text
|
||||
assert "unsharder" in text
|
||||
|
||||
|
||||
def _make_aligner_plan() -> AlignerPlan:
|
||||
unsharder = UnsharderPlan(
|
||||
|
||||
Reference in New Issue
Block a user