Support token aligner planning and execution in dump comparator (#19377)
This commit is contained in:
@@ -119,6 +119,7 @@ class TestExecuteAlignerPlan:
|
||||
x=[self._make_step_plan(step=0, indices=[0, 1])],
|
||||
y=[self._make_step_plan(step=0, indices=[0])],
|
||||
),
|
||||
token_aligner_plan=None,
|
||||
)
|
||||
|
||||
tensors_pair: Pair[list[torch.Tensor]] = Pair(
|
||||
@@ -139,6 +140,7 @@ class TestExecuteAlignerPlan:
|
||||
x=[self._make_step_plan(step=0, indices=[0])],
|
||||
y=[self._make_step_plan(step=0, indices=[0, 1])],
|
||||
),
|
||||
token_aligner_plan=None,
|
||||
)
|
||||
|
||||
tensors_pair: Pair[list[torch.Tensor]] = Pair(
|
||||
@@ -153,12 +155,13 @@ class TestExecuteAlignerPlan:
|
||||
assert result.tensors is None
|
||||
assert result.failed_side_xy == "y"
|
||||
|
||||
def test_single_step(self) -> None:
|
||||
def test_no_token_aligner_single_step(self) -> None:
|
||||
plan = AlignerPlan(
|
||||
per_step_plans=Pair(
|
||||
x=[self._make_step_plan(step=0, indices=[0])],
|
||||
y=[self._make_step_plan(step=0, indices=[0])],
|
||||
),
|
||||
token_aligner_plan=None,
|
||||
)
|
||||
|
||||
t_x: torch.Tensor = torch.tensor([1.0, 2.0])
|
||||
@@ -180,6 +183,7 @@ class TestExecuteAlignerPlan:
|
||||
x=[self._make_step_plan(step=0, indices=[0])],
|
||||
y=[self._make_step_plan(step=0, indices=[0])],
|
||||
),
|
||||
token_aligner_plan=None,
|
||||
)
|
||||
|
||||
tensors_pair: Pair[list[torch.Tensor]] = Pair(
|
||||
|
||||
@@ -136,10 +136,32 @@ class TestComputeAlignerPlan:
|
||||
|
||||
plan: AlignerPlan = compute_aligner_plan(
|
||||
metas_pair=Pair(x=metas_x, y=metas_y),
|
||||
token_aligner_plan=None,
|
||||
)
|
||||
|
||||
assert len(plan.per_step_plans.x) == 1
|
||||
assert len(plan.per_step_plans.y) == 1
|
||||
assert plan.token_aligner_plan is None
|
||||
|
||||
def test_preserves_token_aligner_plan(self) -> None:
|
||||
from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
|
||||
TokenAlignerPlan,
|
||||
TokenLocator,
|
||||
)
|
||||
|
||||
ta_plan = TokenAlignerPlan(
|
||||
locators=Pair(
|
||||
x=TokenLocator(token_index_in_step=[0]),
|
||||
y=TokenLocator(token_index_in_step=[0]),
|
||||
),
|
||||
)
|
||||
|
||||
plan: AlignerPlan = compute_aligner_plan(
|
||||
metas_pair=Pair(x=[_make_meta()], y=[_make_meta()]),
|
||||
token_aligner_plan=ta_plan,
|
||||
)
|
||||
|
||||
assert plan.token_aligner_plan is ta_plan
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user