Support dims annotation and enhance dump loader in dumper (#19276)
This commit is contained in:
@@ -0,0 +1,80 @@
|
|||||||
|
import re
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
|
||||||
|
class ParallelAxis(Enum):
|
||||||
|
TP = "tp"
|
||||||
|
CP = "cp"
|
||||||
|
EP = "ep"
|
||||||
|
SP = "sp"
|
||||||
|
|
||||||
|
|
||||||
|
class Ordering(Enum):
|
||||||
|
ZIGZAG = "zigzag"
|
||||||
|
NATURAL = "natural"
|
||||||
|
|
||||||
|
|
||||||
|
class Reduction(Enum):
|
||||||
|
PARTIAL = "partial"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class DimSpec:
|
||||||
|
name: str
|
||||||
|
parallel: Optional[ParallelAxis] = None
|
||||||
|
ordering: Optional[Ordering] = None
|
||||||
|
reduction: Optional[Reduction] = None
|
||||||
|
|
||||||
|
|
||||||
|
_DIM_PATTERN = re.compile(r"^(?P<name>[a-zA-Z_]\w*)(?:\((?P<modifiers>[^)]+)\))?$")
|
||||||
|
|
||||||
|
_MODIFIER_FIELDS: list[tuple[type[Enum], str]] = [
|
||||||
|
(ParallelAxis, "parallel"),
|
||||||
|
(Ordering, "ordering"),
|
||||||
|
(Reduction, "reduction"),
|
||||||
|
]
|
||||||
|
|
||||||
|
_MODIFIER_LOOKUP: dict[str, tuple[str, Enum]] = {}
|
||||||
|
for _enum_cls, _field in _MODIFIER_FIELDS:
|
||||||
|
for _member in _enum_cls:
|
||||||
|
_MODIFIER_LOOKUP[_member.value] = (_field, _member)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_dim(token: str) -> DimSpec:
|
||||||
|
match = _DIM_PATTERN.match(token)
|
||||||
|
if match is None:
|
||||||
|
raise ValueError(f"Invalid dim token: {token!r}")
|
||||||
|
|
||||||
|
name = match.group("name")
|
||||||
|
modifiers_str = match.group("modifiers")
|
||||||
|
|
||||||
|
if modifiers_str is None:
|
||||||
|
return DimSpec(name=name)
|
||||||
|
|
||||||
|
fields: dict[str, Enum] = {}
|
||||||
|
for part in (p.strip() for p in modifiers_str.split(",")):
|
||||||
|
if part not in _MODIFIER_LOOKUP:
|
||||||
|
raise ValueError(f"Unknown modifier {part!r} in dim spec: {token!r}")
|
||||||
|
field_name, enum_value = _MODIFIER_LOOKUP[part]
|
||||||
|
if field_name in fields:
|
||||||
|
raise ValueError(f"Multiple {field_name} values in dim token: {token!r}")
|
||||||
|
fields[field_name] = enum_value
|
||||||
|
|
||||||
|
return DimSpec(name=name, **fields)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_dims(dims_str: str) -> list[DimSpec]:
|
||||||
|
"""Parse 'b s(cp,zigzag) h(tp) d' -> list[DimSpec]."""
|
||||||
|
if not dims_str.strip():
|
||||||
|
raise ValueError("dims string must not be empty")
|
||||||
|
|
||||||
|
result = [parse_dim(token) for token in dims_str.strip().split()]
|
||||||
|
|
||||||
|
names = [spec.name for spec in result]
|
||||||
|
if len(names) != len(set(names)):
|
||||||
|
duplicates = sorted({n for n in names if names.count(n) > 1})
|
||||||
|
raise ValueError(f"Duplicate dim names: {duplicates}")
|
||||||
|
|
||||||
|
return result
|
||||||
@@ -1,7 +1,9 @@
|
|||||||
import argparse
|
import argparse
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
import polars as pl
|
import polars as pl
|
||||||
|
import torch
|
||||||
|
|
||||||
from sglang.srt.debug_utils.comparator.output_types import (
|
from sglang.srt.debug_utils.comparator.output_types import (
|
||||||
ComparisonRecord,
|
ComparisonRecord,
|
||||||
@@ -11,8 +13,7 @@ from sglang.srt.debug_utils.comparator.output_types import (
|
|||||||
print_record,
|
print_record,
|
||||||
)
|
)
|
||||||
from sglang.srt.debug_utils.comparator.tensor_comparison import compare_tensors
|
from sglang.srt.debug_utils.comparator.tensor_comparison import compare_tensors
|
||||||
from sglang.srt.debug_utils.comparator.utils import load_object
|
from sglang.srt.debug_utils.dump_loader import ValueWithMeta, find_row, read_meta
|
||||||
from sglang.srt.debug_utils.dump_loader import find_row, read_meta
|
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
def main() -> None:
|
||||||
@@ -70,8 +71,8 @@ def run(args: argparse.Namespace) -> None:
|
|||||||
|
|
||||||
path_baseline = Path(args.baseline_path) / row_baseline["filename"]
|
path_baseline = Path(args.baseline_path) / row_baseline["filename"]
|
||||||
|
|
||||||
x_baseline = load_object(path_baseline)
|
x_baseline = _load_tensor(path_baseline)
|
||||||
x_target = load_object(path_target)
|
x_target = _load_tensor(path_target)
|
||||||
|
|
||||||
if x_baseline is None or x_target is None:
|
if x_baseline is None or x_target is None:
|
||||||
counts["skipped"] += 1
|
counts["skipped"] += 1
|
||||||
@@ -104,6 +105,13 @@ def run(args: argparse.Namespace) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _load_tensor(path: Path) -> Optional[torch.Tensor]:
|
||||||
|
loaded = ValueWithMeta.load(path)
|
||||||
|
if not isinstance(loaded.value, torch.Tensor):
|
||||||
|
return None
|
||||||
|
return loaded.value
|
||||||
|
|
||||||
|
|
||||||
def _parse_args() -> argparse.Namespace:
|
def _parse_args() -> argparse.Namespace:
|
||||||
# python -m sglang.srt.debug_utils.comparator --baseline-path ... --target-path ...
|
# python -m sglang.srt.debug_utils.comparator --baseline-path ... --target-path ...
|
||||||
parser = argparse.ArgumentParser()
|
parser = argparse.ArgumentParser()
|
||||||
|
|||||||
@@ -1,5 +1,4 @@
|
|||||||
import functools
|
import functools
|
||||||
from pathlib import Path
|
|
||||||
from typing import Optional, Tuple
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -42,19 +41,3 @@ def calc_rel_diff(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
|
|||||||
denominator = (x * x + y * y).sum()
|
denominator = (x * x + y * y).sum()
|
||||||
sim = 2 * (x * y).sum() / denominator
|
sim = 2 * (x * y).sum() / denominator
|
||||||
return 1 - sim
|
return 1 - sim
|
||||||
|
|
||||||
|
|
||||||
def load_object(path: Path) -> Optional[torch.Tensor]:
|
|
||||||
try:
|
|
||||||
x = torch.load(path, weights_only=False)
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Skip load {path} since error {e}")
|
|
||||||
return None
|
|
||||||
|
|
||||||
if isinstance(x, dict) and "value" in x:
|
|
||||||
x = x["value"]
|
|
||||||
|
|
||||||
if not isinstance(x, torch.Tensor):
|
|
||||||
print(f"Skip load {path} since {type(x)=} is not a Tensor ({x=})")
|
|
||||||
return None
|
|
||||||
return x.cuda()
|
|
||||||
|
|||||||
@@ -1,11 +1,60 @@
|
|||||||
import functools
|
import functools
|
||||||
import os
|
import os
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict
|
from typing import Any, Dict, Tuple
|
||||||
|
|
||||||
import polars as pl
|
import polars as pl
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
_TYPED_FIELDS: list[tuple[str, type]] = [("rank", int)]
|
||||||
|
|
||||||
|
|
||||||
|
def parse_meta_from_filename(path: Path) -> Dict[str, Any]:
|
||||||
|
stem = Path(path).stem
|
||||||
|
result: Dict[str, Any] = {}
|
||||||
|
for kv in stem.split("___"):
|
||||||
|
if "=" in kv:
|
||||||
|
k, v = kv.split("=", 1)
|
||||||
|
result[k] = v
|
||||||
|
for field, converter in _TYPED_FIELDS:
|
||||||
|
if field in result:
|
||||||
|
result[field] = converter(result[field])
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ValueWithMeta:
|
||||||
|
value: Any
|
||||||
|
meta: Dict[str, Any]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def load(path: Path) -> "ValueWithMeta":
|
||||||
|
path = Path(path)
|
||||||
|
meta_from_filename = parse_meta_from_filename(path)
|
||||||
|
|
||||||
|
try:
|
||||||
|
raw = torch.load(path, weights_only=False, map_location="cpu")
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Skip load {path} since error {e}")
|
||||||
|
return ValueWithMeta(
|
||||||
|
value=None, meta={**meta_from_filename, "filename": path.name}
|
||||||
|
)
|
||||||
|
|
||||||
|
value, meta_from_embedded = _unwrap_dict_format(raw)
|
||||||
|
return ValueWithMeta(
|
||||||
|
value=value,
|
||||||
|
meta={**meta_from_filename, **meta_from_embedded, "filename": path.name},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _unwrap_dict_format(obj: Any) -> Tuple[Any, Dict[str, Any]]:
|
||||||
|
if isinstance(obj, dict) and "value" in obj:
|
||||||
|
meta = obj.get("meta", {})
|
||||||
|
assert isinstance(meta, dict), f"Expected meta to be dict, got {type(meta)}"
|
||||||
|
return obj["value"], meta
|
||||||
|
return obj, {}
|
||||||
|
|
||||||
|
|
||||||
class DumpLoader:
|
class DumpLoader:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
@@ -50,10 +99,7 @@ def read_meta(directory):
|
|||||||
rows = []
|
rows = []
|
||||||
for p in directory.glob("*.pt"):
|
for p in directory.glob("*.pt"):
|
||||||
try:
|
try:
|
||||||
full_kwargs = {}
|
full_kwargs = parse_meta_from_filename(p)
|
||||||
for kv in p.stem.split("___"):
|
|
||||||
k, v = kv.split("=")
|
|
||||||
full_kwargs[k] = v
|
|
||||||
rows.append(
|
rows.append(
|
||||||
{
|
{
|
||||||
"filename": str(p.name),
|
"filename": str(p.name),
|
||||||
@@ -83,26 +129,27 @@ def _add_duplicate_index(df: pl.DataFrame) -> pl.DataFrame:
|
|||||||
return df
|
return df
|
||||||
|
|
||||||
|
|
||||||
def find_row(df, conditions: Dict[str, Any]):
|
def filter_rows(df: pl.DataFrame, conditions: Dict[str, Any]) -> list[dict]:
|
||||||
df_sub = df.filter(
|
filter_exprs = [
|
||||||
functools.reduce(
|
|
||||||
lambda a, b: a & b,
|
|
||||||
[
|
|
||||||
(
|
(
|
||||||
pl.col(col)
|
pl.col(col) == _cast_to_polars_dtype(conditions[col], df.schema[col])
|
||||||
== _cast_to_polars_dtype(conditions[col], df.schema[col])
|
|
||||||
if conditions[col] is not None
|
if conditions[col] is not None
|
||||||
else pl.col(col).is_null()
|
else pl.col(col).is_null()
|
||||||
)
|
)
|
||||||
for col in conditions.keys()
|
for col in conditions
|
||||||
if col in df.columns
|
if col in df.columns
|
||||||
],
|
]
|
||||||
)
|
if not filter_exprs:
|
||||||
)
|
return []
|
||||||
if len(df_sub) > 1:
|
return df.filter(functools.reduce(lambda a, b: a & b, filter_exprs)).to_dicts()
|
||||||
print(f"find_row find ambiguous results: {df_sub=}")
|
|
||||||
|
|
||||||
|
def find_row(df: pl.DataFrame, conditions: Dict[str, Any]):
|
||||||
|
rows = filter_rows(df, conditions)
|
||||||
|
if len(rows) > 1:
|
||||||
|
print(f"find_row find ambiguous results: {rows=}")
|
||||||
return None
|
return None
|
||||||
return df_sub.to_dicts()[0] if len(df_sub) > 0 else None
|
return rows[0] if rows else None
|
||||||
|
|
||||||
|
|
||||||
def _cast_to_polars_dtype(value, target_dtype):
|
def _cast_to_polars_dtype(value, target_dtype):
|
||||||
|
|||||||
@@ -223,7 +223,24 @@ class _Dumper:
|
|||||||
self._state.step += 1
|
self._state.step += 1
|
||||||
print(f"[Dumper] [{time.time()}] step={self._state.step}")
|
print(f"[Dumper] [{time.time()}] step={self._state.step}")
|
||||||
|
|
||||||
def dump(self, name: str, value, save: bool = True, **kwargs) -> None:
|
def dump(
|
||||||
|
self,
|
||||||
|
name: str,
|
||||||
|
value,
|
||||||
|
save: bool = True,
|
||||||
|
dims: Optional[str] = None,
|
||||||
|
dims_grad: Optional[str] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
|
value_meta: dict = {}
|
||||||
|
grad_meta: dict = {}
|
||||||
|
if dims is not None:
|
||||||
|
value_meta["dims"] = dims
|
||||||
|
grad_meta["dims"] = dims
|
||||||
|
if dims_grad is not None:
|
||||||
|
value_meta["dims_grad"] = dims_grad
|
||||||
|
grad_meta["dims"] = dims_grad
|
||||||
|
|
||||||
self._dump_inner(
|
self._dump_inner(
|
||||||
name=name,
|
name=name,
|
||||||
value=value,
|
value=value,
|
||||||
@@ -234,6 +251,8 @@ class _Dumper:
|
|||||||
enable_future_grad=self._config.enable_grad,
|
enable_future_grad=self._config.enable_grad,
|
||||||
value_tag="Dumper.Value",
|
value_tag="Dumper.Value",
|
||||||
grad_tag="Dumper.Grad",
|
grad_tag="Dumper.Grad",
|
||||||
|
value_meta_only_fields=value_meta,
|
||||||
|
grad_meta_only_fields=grad_meta,
|
||||||
)
|
)
|
||||||
|
|
||||||
def dump_model(
|
def dump_model(
|
||||||
@@ -336,6 +355,8 @@ class _Dumper:
|
|||||||
enable_future_grad: bool,
|
enable_future_grad: bool,
|
||||||
value_tag: str,
|
value_tag: str,
|
||||||
grad_tag: str,
|
grad_tag: str,
|
||||||
|
value_meta_only_fields: Optional[dict] = None,
|
||||||
|
grad_meta_only_fields: Optional[dict] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._http_manager # noqa: B018
|
self._http_manager # noqa: B018
|
||||||
|
|
||||||
@@ -359,6 +380,7 @@ class _Dumper:
|
|||||||
tags=tags,
|
tags=tags,
|
||||||
value=value,
|
value=value,
|
||||||
save=save,
|
save=save,
|
||||||
|
meta_only_fields=value_meta_only_fields or {},
|
||||||
)
|
)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
@@ -371,6 +393,7 @@ class _Dumper:
|
|||||||
tags={**tags, "name": f"grad__{name}"},
|
tags={**tags, "name": f"grad__{name}"},
|
||||||
value=g,
|
value=g,
|
||||||
save=save,
|
save=save,
|
||||||
|
meta_only_fields=grad_meta_only_fields or {},
|
||||||
)
|
)
|
||||||
|
|
||||||
if enable_future_grad:
|
if enable_future_grad:
|
||||||
@@ -379,6 +402,7 @@ class _Dumper:
|
|||||||
tensor=value,
|
tensor=value,
|
||||||
extra_kwargs=extra_kwargs,
|
extra_kwargs=extra_kwargs,
|
||||||
save=save,
|
save=save,
|
||||||
|
meta_only_fields=grad_meta_only_fields or {},
|
||||||
)
|
)
|
||||||
|
|
||||||
def _register_dump_grad_hook(
|
def _register_dump_grad_hook(
|
||||||
@@ -388,6 +412,7 @@ class _Dumper:
|
|||||||
tensor,
|
tensor,
|
||||||
extra_kwargs: dict,
|
extra_kwargs: dict,
|
||||||
save: bool,
|
save: bool,
|
||||||
|
meta_only_fields: Optional[dict] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
if not isinstance(tensor, torch.Tensor):
|
if not isinstance(tensor, torch.Tensor):
|
||||||
return
|
return
|
||||||
@@ -396,6 +421,7 @@ class _Dumper:
|
|||||||
|
|
||||||
captured_step = self._state.step
|
captured_step = self._state.step
|
||||||
captured_tags = dict(name=f"grad__{name}", **deepcopy(extra_kwargs))
|
captured_tags = dict(name=f"grad__{name}", **deepcopy(extra_kwargs))
|
||||||
|
captured_meta_only = meta_only_fields or {}
|
||||||
|
|
||||||
def grad_hook(grad: torch.Tensor) -> None:
|
def grad_hook(grad: torch.Tensor) -> None:
|
||||||
self._dump_single(
|
self._dump_single(
|
||||||
@@ -404,6 +430,7 @@ class _Dumper:
|
|||||||
value=grad,
|
value=grad,
|
||||||
save=save,
|
save=save,
|
||||||
step=captured_step,
|
step=captured_step,
|
||||||
|
meta_only_fields=captured_meta_only,
|
||||||
)
|
)
|
||||||
|
|
||||||
tensor.register_hook(grad_hook)
|
tensor.register_hook(grad_hook)
|
||||||
@@ -416,6 +443,7 @@ class _Dumper:
|
|||||||
value,
|
value,
|
||||||
save: bool,
|
save: bool,
|
||||||
step: Optional[int] = None,
|
step: Optional[int] = None,
|
||||||
|
meta_only_fields: Optional[dict] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._ensure_exp_name()
|
self._ensure_exp_name()
|
||||||
self._state.dump_index += 1
|
self._state.dump_index += 1
|
||||||
@@ -445,7 +473,11 @@ class _Dumper:
|
|||||||
if save and (self._config.enable_output_file or capturing):
|
if save and (self._config.enable_output_file or capturing):
|
||||||
output_data = {
|
output_data = {
|
||||||
"value": value,
|
"value": value,
|
||||||
"meta": dict(**full_kwargs, **self._static_meta),
|
"meta": dict(
|
||||||
|
**full_kwargs,
|
||||||
|
**self._static_meta,
|
||||||
|
**(meta_only_fields or {}),
|
||||||
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
if capturing:
|
if capturing:
|
||||||
|
|||||||
@@ -0,0 +1,101 @@
|
|||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.dims import (
|
||||||
|
DimSpec,
|
||||||
|
Ordering,
|
||||||
|
ParallelAxis,
|
||||||
|
Reduction,
|
||||||
|
parse_dim,
|
||||||
|
parse_dims,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=10, suite="default", nightly=True)
|
||||||
|
|
||||||
|
|
||||||
|
class TestParseDim:
|
||||||
|
def test_plain_name(self) -> None:
|
||||||
|
assert parse_dim("b") == DimSpec(name="b")
|
||||||
|
|
||||||
|
def test_parallel_axis(self) -> None:
|
||||||
|
assert parse_dim("h(tp)") == DimSpec(name="h", parallel=ParallelAxis.TP)
|
||||||
|
|
||||||
|
def test_all_parallel_axes(self) -> None:
|
||||||
|
assert parse_dim("a(tp)").parallel == ParallelAxis.TP
|
||||||
|
assert parse_dim("a(cp)").parallel == ParallelAxis.CP
|
||||||
|
assert parse_dim("a(ep)").parallel == ParallelAxis.EP
|
||||||
|
assert parse_dim("a(sp)").parallel == ParallelAxis.SP
|
||||||
|
|
||||||
|
def test_ordering(self) -> None:
|
||||||
|
assert parse_dim("s(cp,zigzag)").ordering == Ordering.ZIGZAG
|
||||||
|
assert parse_dim("s(cp,natural)").ordering == Ordering.NATURAL
|
||||||
|
|
||||||
|
def test_reduction(self) -> None:
|
||||||
|
assert parse_dim("h(tp,partial)").reduction == Reduction.PARTIAL
|
||||||
|
|
||||||
|
def test_all_modifiers(self) -> None:
|
||||||
|
assert parse_dim("s(cp,zigzag,partial)") == DimSpec(
|
||||||
|
name="s",
|
||||||
|
parallel=ParallelAxis.CP,
|
||||||
|
ordering=Ordering.ZIGZAG,
|
||||||
|
reduction=Reduction.PARTIAL,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_invalid_token_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="Invalid dim token"):
|
||||||
|
parse_dim("h()")
|
||||||
|
with pytest.raises(ValueError, match="Invalid dim token"):
|
||||||
|
parse_dim("h(tp(x))")
|
||||||
|
|
||||||
|
def test_unknown_modifier_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="Unknown modifier"):
|
||||||
|
parse_dim("h(xyz)")
|
||||||
|
with pytest.raises(ValueError, match="Unknown modifier"):
|
||||||
|
parse_dim("h(tp,foobar)")
|
||||||
|
|
||||||
|
def test_multiple_ordering_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="Multiple ordering"):
|
||||||
|
parse_dim("s(cp,zigzag,natural)")
|
||||||
|
|
||||||
|
def test_multiple_reduction_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="Multiple reduction"):
|
||||||
|
parse_dim("h(tp,partial,partial)")
|
||||||
|
|
||||||
|
|
||||||
|
class TestParseDims:
|
||||||
|
def test_multi_dims(self) -> None:
|
||||||
|
assert parse_dims("b s h d") == [
|
||||||
|
DimSpec(name="b"),
|
||||||
|
DimSpec(name="s"),
|
||||||
|
DimSpec(name="h"),
|
||||||
|
DimSpec(name="d"),
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_single_dim(self) -> None:
|
||||||
|
assert parse_dims("t") == [DimSpec(name="t")]
|
||||||
|
|
||||||
|
def test_mixed_annotated(self) -> None:
|
||||||
|
assert parse_dims("b s(cp,zigzag) h(tp) d") == [
|
||||||
|
DimSpec(name="b"),
|
||||||
|
DimSpec(name="s", parallel=ParallelAxis.CP, ordering=Ordering.ZIGZAG),
|
||||||
|
DimSpec(name="h", parallel=ParallelAxis.TP),
|
||||||
|
DimSpec(name="d"),
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_empty_string_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="empty"):
|
||||||
|
parse_dims("")
|
||||||
|
|
||||||
|
def test_whitespace_only_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="empty"):
|
||||||
|
parse_dims(" ")
|
||||||
|
|
||||||
|
def test_duplicate_name_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="Duplicate"):
|
||||||
|
parse_dims("h h")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__]))
|
||||||
@@ -1,5 +1,4 @@
|
|||||||
import sys
|
import sys
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
@@ -8,7 +7,6 @@ from sglang.srt.debug_utils.comparator.utils import (
|
|||||||
argmax_coord,
|
argmax_coord,
|
||||||
calc_rel_diff,
|
calc_rel_diff,
|
||||||
compute_smaller_dtype,
|
compute_smaller_dtype,
|
||||||
load_object,
|
|
||||||
try_unify_shape,
|
try_unify_shape,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -92,28 +90,5 @@ class TestComputeSmallerDtype:
|
|||||||
assert compute_smaller_dtype(torch.int32, torch.int64) is None
|
assert compute_smaller_dtype(torch.int32, torch.int64) is None
|
||||||
|
|
||||||
|
|
||||||
class TestLoadObject:
|
|
||||||
def test_load_tensor(self, tmp_path):
|
|
||||||
path = tmp_path / "tensor.pt"
|
|
||||||
torch.save(torch.randn(5, 5), path)
|
|
||||||
assert load_object(path).shape == (5, 5)
|
|
||||||
|
|
||||||
def test_load_dict_with_value_key(self, tmp_path):
|
|
||||||
path = tmp_path / "wrapped.pt"
|
|
||||||
tensor = torch.randn(3, 3)
|
|
||||||
torch.save({"value": tensor}, path)
|
|
||||||
result = load_object(path)
|
|
||||||
assert result is not None
|
|
||||||
assert result.shape == (3, 3)
|
|
||||||
|
|
||||||
def test_non_tensor_returns_none(self, tmp_path):
|
|
||||||
path = tmp_path / "tensor.pt"
|
|
||||||
torch.save({"dict": 1}, path)
|
|
||||||
assert load_object(path) is None
|
|
||||||
|
|
||||||
def test_nonexistent_returns_none(self):
|
|
||||||
assert load_object(Path("/nonexistent.pt")) is None
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
@@ -1,50 +1,58 @@
|
|||||||
import tempfile
|
import sys
|
||||||
import unittest
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import polars as pl
|
import polars as pl
|
||||||
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.dump_loader import (
|
||||||
|
ValueWithMeta,
|
||||||
|
_add_duplicate_index,
|
||||||
|
_cast_to_polars_dtype,
|
||||||
|
find_row,
|
||||||
|
read_meta,
|
||||||
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
|
||||||
|
|
||||||
register_cpu_ci(est_time=30, suite="default", nightly=True)
|
register_cpu_ci(est_time=30, suite="default", nightly=True)
|
||||||
|
|
||||||
|
|
||||||
class TestDumpLoader(CustomTestCase):
|
class TestReadMeta:
|
||||||
def test_read_meta(self):
|
def test_basic(self, tmp_path):
|
||||||
from sglang.srt.debug_utils.dump_loader import read_meta
|
|
||||||
|
|
||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
|
||||||
for fn in [
|
for fn in [
|
||||||
"step=1___rank=0___dump_index=1___name=a.pt",
|
"step=1___rank=0___dump_index=1___name=a.pt",
|
||||||
"step=2___rank=0___dump_index=2___name=b.pt",
|
"step=2___rank=0___dump_index=2___name=b.pt",
|
||||||
]:
|
]:
|
||||||
torch.save(torch.randn(5), Path(tmpdir) / fn)
|
torch.save(torch.randn(5), tmp_path / fn)
|
||||||
|
|
||||||
df = read_meta(tmpdir)
|
df = read_meta(str(tmp_path))
|
||||||
self.assertEqual(len(df), 2)
|
assert len(df) == 2
|
||||||
self.assertTrue(all(c in df.columns for c in ["step", "rank", "name"]))
|
assert all(c in df.columns for c in ["step", "rank", "name"])
|
||||||
|
|
||||||
def test_find_row(self):
|
|
||||||
from sglang.srt.debug_utils.dump_loader import find_row
|
|
||||||
|
|
||||||
|
class TestFindRow:
|
||||||
|
def test_single_match(self):
|
||||||
df = pl.DataFrame({"id": [1, 2], "name": ["a", "b"], "file": ["f1", "f2"]})
|
df = pl.DataFrame({"id": [1, 2], "name": ["a", "b"], "file": ["f1", "f2"]})
|
||||||
self.assertEqual(find_row(df, {"id": 2})["file"], "f2")
|
assert find_row(df, {"id": 2})["file"] == "f2"
|
||||||
self.assertIsNone(find_row(df, {"id": 999}))
|
|
||||||
|
|
||||||
df_dup = pl.DataFrame({"id": [1, 1], "file": ["f1", "f2"]})
|
def test_no_match(self):
|
||||||
self.assertIsNone(find_row(df_dup, {"id": 1}))
|
df = pl.DataFrame({"id": [1, 2], "name": ["a", "b"], "file": ["f1", "f2"]})
|
||||||
|
assert find_row(df, {"id": 999}) is None
|
||||||
|
|
||||||
def test_cast_to_polars_dtype(self):
|
def test_ambiguous(self):
|
||||||
from sglang.srt.debug_utils.dump_loader import _cast_to_polars_dtype
|
df = pl.DataFrame({"id": [1, 1], "file": ["f1", "f2"]})
|
||||||
|
assert find_row(df, {"id": 1}) is None
|
||||||
|
|
||||||
self.assertEqual(_cast_to_polars_dtype("42", pl.Int64), 42)
|
|
||||||
self.assertEqual(_cast_to_polars_dtype("3.14", pl.Float64), 3.14)
|
|
||||||
|
|
||||||
def test_add_duplicate_index(self):
|
class TestCastToPolars:
|
||||||
from sglang.srt.debug_utils.dump_loader import _add_duplicate_index
|
def test_int(self):
|
||||||
|
assert _cast_to_polars_dtype("42", pl.Int64) == 42
|
||||||
|
|
||||||
|
def test_float(self):
|
||||||
|
assert _cast_to_polars_dtype("3.14", pl.Float64) == pytest.approx(3.14)
|
||||||
|
|
||||||
|
|
||||||
|
class TestAddDuplicateIndex:
|
||||||
|
def test_basic(self):
|
||||||
df = pl.DataFrame(
|
df = pl.DataFrame(
|
||||||
{
|
{
|
||||||
"name": ["a", "a", "b"],
|
"name": ["a", "a", "b"],
|
||||||
@@ -53,13 +61,40 @@ class TestDumpLoader(CustomTestCase):
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
result = _add_duplicate_index(df)
|
result = _add_duplicate_index(df)
|
||||||
self.assertEqual(
|
assert result.filter(pl.col("name") == "a").sort("dump_index")[
|
||||||
result.filter(pl.col("name") == "a")
|
"duplicate_index"
|
||||||
.sort("dump_index")["duplicate_index"]
|
].to_list() == [0, 1]
|
||||||
.to_list(),
|
|
||||||
[0, 1],
|
|
||||||
)
|
class TestValueWithMeta:
|
||||||
|
def test_load_dict_format(self, tmp_path) -> None:
|
||||||
|
path = tmp_path / "step=0___rank=0___dump_index=1___name=hidden.pt"
|
||||||
|
tensor = torch.randn(4, 8)
|
||||||
|
torch.save({"value": tensor, "meta": {"custom": "field"}}, path)
|
||||||
|
|
||||||
|
loaded = ValueWithMeta.load(path)
|
||||||
|
assert torch.allclose(loaded.value, tensor)
|
||||||
|
assert loaded.meta["custom"] == "field"
|
||||||
|
assert loaded.meta["name"] == "hidden"
|
||||||
|
assert loaded.meta["rank"] == 0
|
||||||
|
|
||||||
|
def test_load_bare_tensor(self, tmp_path) -> None:
|
||||||
|
path = tmp_path / "step=0___rank=0___dump_index=1___name=bare.pt"
|
||||||
|
tensor = torch.randn(3, 3)
|
||||||
|
torch.save(tensor, path)
|
||||||
|
|
||||||
|
loaded = ValueWithMeta.load(path)
|
||||||
|
assert torch.allclose(loaded.value, tensor)
|
||||||
|
assert loaded.meta["name"] == "bare"
|
||||||
|
|
||||||
|
def test_load_corrupted_file(self, tmp_path) -> None:
|
||||||
|
path = tmp_path / "step=0___rank=0___dump_index=1___name=bad.pt"
|
||||||
|
path.write_text("not a valid pt file")
|
||||||
|
|
||||||
|
loaded = ValueWithMeta.load(path)
|
||||||
|
assert loaded.value is None
|
||||||
|
assert loaded.meta["name"] == "bad"
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
@@ -2256,5 +2256,74 @@ class TestNonIntrusiveKwargsModel(_NonIntrusiveTestBase):
|
|||||||
assert captured["qkv_format"]["value"] == "thd"
|
assert captured["qkv_format"]["value"] == "thd"
|
||||||
|
|
||||||
|
|
||||||
|
class TestDumperDims:
|
||||||
|
def test_dims_in_meta_not_filename(self, tmp_path) -> None:
|
||||||
|
dumper = _make_test_dumper(tmp_path)
|
||||||
|
tensor = torch.randn(4, 8)
|
||||||
|
dumper.dump("hidden", tensor, dims="b h(tp)")
|
||||||
|
dumper.step()
|
||||||
|
|
||||||
|
exp_dir = tmp_path / dumper._config.exp_name
|
||||||
|
pt_files = list(exp_dir.glob("*.pt"))
|
||||||
|
assert len(pt_files) == 1
|
||||||
|
|
||||||
|
assert "dims" not in pt_files[0].stem
|
||||||
|
|
||||||
|
data = torch.load(pt_files[0], weights_only=False)
|
||||||
|
assert "dims" in data["meta"]
|
||||||
|
assert data["meta"]["dims"] == "b h(tp)"
|
||||||
|
|
||||||
|
def test_dims_grad_override(self, tmp_path) -> None:
|
||||||
|
dumper = _Dumper(
|
||||||
|
config=DumperConfig(
|
||||||
|
enable=True,
|
||||||
|
dir=str(tmp_path),
|
||||||
|
enable_http_server=False,
|
||||||
|
enable_grad=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
tensor = torch.randn(4, 8, requires_grad=True)
|
||||||
|
dumper.dump("hidden", tensor, dims="b h(tp)", dims_grad="b h(tp,partial)")
|
||||||
|
dumper.step()
|
||||||
|
|
||||||
|
tensor.backward(torch.ones_like(tensor))
|
||||||
|
|
||||||
|
exp_dir = tmp_path / dumper._config.exp_name
|
||||||
|
pt_files = sorted(exp_dir.glob("*.pt"))
|
||||||
|
assert len(pt_files) == 2
|
||||||
|
|
||||||
|
value_file = [f for f in pt_files if "grad__" not in f.stem][0]
|
||||||
|
grad_file = [f for f in pt_files if "grad__" in f.stem][0]
|
||||||
|
|
||||||
|
value_data = torch.load(value_file, weights_only=False)
|
||||||
|
assert value_data["meta"]["dims"] == "b h(tp)"
|
||||||
|
assert value_data["meta"]["dims_grad"] == "b h(tp,partial)"
|
||||||
|
|
||||||
|
grad_data = torch.load(grad_file, weights_only=False)
|
||||||
|
assert grad_data["meta"]["dims"] == "b h(tp,partial)"
|
||||||
|
|
||||||
|
def test_dims_grad_inherits(self, tmp_path) -> None:
|
||||||
|
dumper = _Dumper(
|
||||||
|
config=DumperConfig(
|
||||||
|
enable=True,
|
||||||
|
dir=str(tmp_path),
|
||||||
|
enable_http_server=False,
|
||||||
|
enable_grad=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
tensor = torch.randn(4, 8, requires_grad=True)
|
||||||
|
dumper.dump("hidden", tensor, dims="b h(tp)")
|
||||||
|
dumper.step()
|
||||||
|
|
||||||
|
tensor.backward(torch.ones_like(tensor))
|
||||||
|
|
||||||
|
exp_dir = tmp_path / dumper._config.exp_name
|
||||||
|
grad_file = [f for f in exp_dir.glob("*.pt") if "grad__" in f.stem][0]
|
||||||
|
grad_data = torch.load(grad_file, weights_only=False)
|
||||||
|
assert grad_data["meta"]["dims"] == "b h(tp)"
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
Reference in New Issue
Block a user