Trace execution information in dump comparator (#19682)

This commit is contained in:
fzyzcjy
2026-03-02 18:46:27 +08:00
committed by GitHub
parent 3ebd85bf1c
commit 5bf3deb4bc
17 changed files with 1180 additions and 148 deletions

View File

@@ -5,6 +5,8 @@ import torch
from sglang.srt.debug_utils.comparator.aligner.entrypoint.executor import (
AlignerResult,
StepPlansResult,
SubPlansResult,
_execute_step_plans,
execute_aligner_plan,
execute_sub_plan,
@@ -31,25 +33,28 @@ register_cpu_ci(est_time=15, suite="default", nightly=True)
class TestExecuteSubPlans:
def test_empty_tensors_returns_none(self) -> None:
result, checks = execute_sub_plans(tensors=[], plans=[])
assert result is None
assert checks == []
r: SubPlansResult = execute_sub_plans(tensors=[], plans=[])
assert r.tensor is None
assert r.checks == []
assert r.snapshots == []
def test_no_plans_single_tensor_passthrough(self) -> None:
tensor: torch.Tensor = torch.tensor([1.0, 2.0, 3.0])
result, checks = execute_sub_plans(tensors=[tensor], plans=[])
assert result is not None
assert torch.equal(result, tensor)
assert checks == []
r: SubPlansResult = execute_sub_plans(tensors=[tensor], plans=[])
assert r.tensor is not None
assert torch.equal(r.tensor, tensor)
assert r.checks == []
assert r.snapshots == []
def test_no_plans_multiple_tensors_returns_none(self) -> None:
tensors: list[torch.Tensor] = [
torch.tensor([1.0]),
torch.tensor([2.0]),
]
result, checks = execute_sub_plans(tensors=tensors, plans=[])
assert result is None
assert checks == []
r: SubPlansResult = execute_sub_plans(tensors=tensors, plans=[])
assert r.tensor is None
assert r.checks == []
assert r.snapshots == []
def test_with_unsharder_plan(self) -> None:
t0: torch.Tensor = torch.tensor([[1.0, 2.0]]).refine_names("b", "h")
@@ -61,12 +66,13 @@ class TestExecuteSubPlans:
groups=[[0, 1]],
)
result, checks = execute_sub_plans(tensors=[t0, t1], plans=[plan])
r: SubPlansResult = execute_sub_plans(tensors=[t0, t1], plans=[plan])
assert result is not None
assert r.tensor is not None
expected: torch.Tensor = torch.tensor([[1.0, 2.0, 3.0, 4.0]])
assert torch.equal(result.rename(None), expected)
assert checks == []
assert torch.equal(r.tensor.rename(None), expected)
assert r.checks == []
assert len(r.snapshots) == 1
class TestExecuteSubPlan:
@@ -91,10 +97,13 @@ class TestExecuteStepPlans:
sub_plans=[],
)
result, checks = _execute_step_plans(tensors=tensors, step_plans=[step_plan])
r: StepPlansResult = _execute_step_plans(
tensors=tensors, step_plans=[step_plan]
)
assert result == {}
assert checks == []
assert r.tensors == {}
assert r.checks == []
assert len(r.traced_side.step_plans) == 1
def test_single_step_passthrough(self) -> None:
tensor: torch.Tensor = torch.tensor([1.0, 2.0])
@@ -105,11 +114,15 @@ class TestExecuteStepPlans:
sub_plans=[],
)
result, checks = _execute_step_plans(tensors=[tensor], step_plans=[step_plan])
r: StepPlansResult = _execute_step_plans(
tensors=[tensor], step_plans=[step_plan]
)
assert 5 in result
assert torch.equal(result[5], tensor)
assert checks == []
assert 5 in r.tensors
assert torch.equal(r.tensors[5], tensor)
assert r.checks == []
assert len(r.traced_side.step_plans) == 1
assert r.traced_side.step_plans[0].step == 5
class TestExecuteAlignerPlan: