Purge usage of pytorch named tensors (#25911)
This commit is contained in:
@@ -7,9 +7,10 @@ from sglang.srt.debug_utils.comparator.dims_spec import (
|
||||
DimSpec,
|
||||
apply_dim_names,
|
||||
find_dim_index,
|
||||
get_dim_names,
|
||||
parse_dims,
|
||||
resolve_dim_by_name,
|
||||
strip_dim_names,
|
||||
without_dim_names,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
@@ -43,13 +44,13 @@ class TestFindDimIndex:
|
||||
|
||||
class TestResolveDimByName:
|
||||
def test_resolve_found(self) -> None:
|
||||
tensor: torch.Tensor = torch.randn(2, 3, 4).refine_names("b", "s", "h")
|
||||
tensor: torch.Tensor = apply_dim_names(torch.randn(2, 3, 4), ["b", "s", "h"])
|
||||
assert resolve_dim_by_name(tensor, "b") == 0
|
||||
assert resolve_dim_by_name(tensor, "s") == 1
|
||||
assert resolve_dim_by_name(tensor, "h") == 2
|
||||
|
||||
def test_resolve_not_found_raises(self) -> None:
|
||||
tensor: torch.Tensor = torch.randn(2, 3).refine_names("b", "s")
|
||||
tensor: torch.Tensor = apply_dim_names(torch.randn(2, 3), ["b", "s"])
|
||||
with pytest.raises(ValueError, match="not in tensor names"):
|
||||
resolve_dim_by_name(tensor, "h")
|
||||
|
||||
@@ -63,13 +64,13 @@ class TestApplyDimNames:
|
||||
def test_apply(self) -> None:
|
||||
tensor: torch.Tensor = torch.randn(2, 3, 4)
|
||||
named: torch.Tensor = apply_dim_names(tensor, ["b", "s", "h"])
|
||||
assert named.names == ("b", "s", "h")
|
||||
assert get_dim_names(named) == ("b", "s", "h")
|
||||
assert named.shape == (2, 3, 4)
|
||||
|
||||
def test_apply_preserves_data(self) -> None:
|
||||
tensor: torch.Tensor = torch.randn(2, 3)
|
||||
named: torch.Tensor = apply_dim_names(tensor, ["x", "y"])
|
||||
assert torch.equal(strip_dim_names(named), tensor)
|
||||
assert torch.equal(without_dim_names(named), tensor)
|
||||
|
||||
def test_ndim_mismatch_gives_clear_error(self) -> None:
|
||||
tensor: torch.Tensor = torch.randn(10, 1, 128)
|
||||
@@ -82,14 +83,14 @@ class TestApplyDimNames:
|
||||
|
||||
class TestStripDimNames:
|
||||
def test_strip(self) -> None:
|
||||
tensor: torch.Tensor = torch.randn(2, 3).refine_names("a", "b")
|
||||
stripped: torch.Tensor = strip_dim_names(tensor)
|
||||
assert stripped.names == (None, None)
|
||||
tensor: torch.Tensor = apply_dim_names(torch.randn(2, 3), ["a", "b"])
|
||||
stripped: torch.Tensor = without_dim_names(tensor)
|
||||
assert get_dim_names(stripped) == (None, None)
|
||||
|
||||
def test_strip_already_unnamed(self) -> None:
|
||||
tensor: torch.Tensor = torch.randn(2, 3)
|
||||
stripped: torch.Tensor = strip_dim_names(tensor)
|
||||
assert stripped.names == (None, None)
|
||||
stripped: torch.Tensor = without_dim_names(tensor)
|
||||
assert get_dim_names(stripped) == (None, None)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user