Purge usage of pytorch named tensors (#25911)

This commit is contained in:
Joel Schlosser
2026-05-26 14:58:57 -07:00
committed by GitHub
parent 1a05b511e4
commit 6989fede3c
17 changed files with 247 additions and 159 deletions
@@ -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__":