Support loading token aligner data in dump comparator (#19376)

This commit is contained in:
fzyzcjy
2026-02-26 10:03:56 +08:00
committed by GitHub
parent e8dd14519d
commit d34d5aca07
10 changed files with 1182 additions and 43 deletions
@@ -3,6 +3,13 @@ import sys
import pytest
from pydantic import ValidationError
from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
PositionalSeqId,
TokenAlignerPlan,
TokenAlignerSeqInfo,
TokenAlignerStepAux,
TokenLocator,
)
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo
from sglang.srt.debug_utils.comparator.output_types import (
ComparisonRecord,
@@ -15,7 +22,7 @@ from sglang.srt.debug_utils.comparator.tensor_comparator.types import (
TensorInfo,
TensorStats,
)
from sglang.srt.debug_utils.comparator.utils import _check_equal_lengths
from sglang.srt.debug_utils.comparator.utils import Pair, _check_equal_lengths
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="default", nightly=True)
@@ -33,6 +40,109 @@ class TestCheckEqualLengths:
_check_equal_lengths(a=[1, 2], b=[3])
class TestTokenAlignerStepAux:
def test_valid(self):
aux = TokenAlignerStepAux(
input_ids=[10, 20, 30],
positions=[0, 1, 2],
seq_lens=[2, 1],
seq_ids=[
PositionalSeqId(step=0, seq_index=0),
PositionalSeqId(step=0, seq_index=1),
],
)
assert len(aux.input_ids) == 3
def test_token_length_mismatch(self):
with pytest.raises(ValueError, match="Length mismatch"):
TokenAlignerStepAux(
input_ids=[10, 20, 30],
positions=[0, 1],
seq_lens=[2, 1],
seq_ids=[
PositionalSeqId(step=0, seq_index=0),
PositionalSeqId(step=0, seq_index=1),
],
)
def test_seq_length_mismatch(self):
with pytest.raises(ValueError, match="Length mismatch"):
TokenAlignerStepAux(
input_ids=[10, 20, 30],
positions=[0, 1, 2],
seq_lens=[2, 1],
seq_ids=[PositionalSeqId(step=0, seq_index=0)],
)
def test_sum_seq_lens_mismatch(self):
with pytest.raises(ValueError, match="sum\\(seq_lens\\)"):
TokenAlignerStepAux(
input_ids=[10, 20, 30],
positions=[0, 1, 2],
seq_lens=[1, 1],
seq_ids=[
PositionalSeqId(step=0, seq_index=0),
PositionalSeqId(step=0, seq_index=1),
],
)
class TestTokenAlignerSeqInfo:
def test_valid(self):
info = TokenAlignerSeqInfo(
input_ids=[10, 20, 30],
positions=[0, 1, 2],
locator=TokenLocator(token_index_in_step=[0, 1, 0]),
)
assert len(info.input_ids) == 3
def test_length_mismatch(self):
with pytest.raises(ValidationError):
TokenAlignerSeqInfo(
input_ids=[10, 20, 30],
positions=[0, 1, 2],
locator=TokenLocator(token_index_in_step=[0, 1]),
)
def test_positions_not_sequential(self):
with pytest.raises(ValidationError, match="positions must be"):
TokenAlignerSeqInfo(
input_ids=[10, 20, 30],
positions=[0, 2, 1],
locator=TokenLocator(token_index_in_step=[0, 1, 0]),
)
class TestTokenAlignerPlan:
def test_valid(self):
plan = TokenAlignerPlan(
locators=Pair(
x=TokenLocator(token_index_in_step=[0, 1, 0]),
y=TokenLocator(token_index_in_step=[0, 0, 1]),
),
)
assert len(plan.locators.x.token_index_in_step) == 3
def test_length_mismatch(self):
with pytest.raises(ValidationError, match="Length mismatch"):
TokenAlignerPlan(
locators=Pair(
x=TokenLocator(token_index_in_step=[0, 1]),
y=TokenLocator(token_index_in_step=[0, 0, 1]),
),
)
class TestSummaryRecord:
def test_valid(self):
record = SummaryRecord(total=10, passed=7, failed=2, skipped=1)
assert record.total == 10
def test_total_mismatch(self):
with pytest.raises(ValidationError, match="total=10"):
SummaryRecord(total=10, passed=5, failed=2, skipped=1)
class TestAxisInfo:
def test_valid(self):
info = AxisInfo(axis_rank=0, axis_size=4)
@@ -59,16 +169,6 @@ class TestAxisInfo:
assert info.axis_rank == 3
class TestSummaryRecord:
def test_valid(self):
record = SummaryRecord(total=10, passed=7, failed=2, skipped=1)
assert record.total == 10
def test_total_mismatch(self):
with pytest.raises(ValidationError, match="total=10"):
SummaryRecord(total=10, passed=5, failed=2, skipped=1)
def _make_tensor_info() -> TensorInfo:
return TensorInfo(
shape=[4, 4],