Add kv_canary PP self-test fixture and SWA divergence coverage (#27410)

This commit is contained in:
fzyzcjy
2026-06-06 09:06:19 +08:00
committed by GitHub
parent b3e4c204fd
commit 21201ef718
8 changed files with 208 additions and 21 deletions
@@ -0,0 +1,24 @@
from __future__ import annotations
import unittest
from sglang.srt.kv_canary.config import CanaryMode
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kv_canary.pp_fixture import CanaryPPFixture
register_cuda_ci(est_time=220, stage="extra-a", runner_config="2-gpu-large")
class TestPPBaselineSwa(CanaryPPFixture):
kv_canary_mode = CanaryMode.LOG
def test_no_violation(self) -> None:
for _ in range(self.workload_n_batches):
self.send_parallel_requests()
self.assert_no_violation(wait_seconds=2.0)
self.maybe_assert_swa_divergence_observed()
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,40 @@
from __future__ import annotations
import unittest
from typing import ClassVar
from sglang.srt.kv_canary.config import CanaryMode
from sglang.srt.kv_canary.perturb.config import TargetGroupKind
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kv_canary.pp_fixture import CanaryPPFixture
register_cuda_ci(est_time=220, stage="extra-a", runner_config="2-gpu-large")
class TestPPPerturbSwaSwa(CanaryPPFixture):
kv_canary_mode = CanaryMode.LOG
target_group: ClassVar[TargetGroupKind] = TargetGroupKind.SWA
extra_server_args = ("--kv-canary-real-data", "partial")
@classmethod
def setUpClass(cls) -> None:
cls.extra_env = {
"SGLANG_KV_CANARY_PERTURB_REAL_KV_USED_PROB": "0.1",
"SGLANG_KV_CANARY_PERTURB_TARGET_GROUP": str(cls.target_group),
"SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS": "0",
}
super().setUpClass()
def test_real_kv_used_perturbation_reports_real_kv_hash_violation(self) -> None:
for _ in range(self.workload_n_batches):
self.send_parallel_requests()
self.assert_per_forward_violation_reported(
fail_reason="verify_real_kv_hash",
target_group=self.target_group,
)
self.maybe_assert_swa_divergence_observed()
if __name__ == "__main__":
unittest.main()
@@ -69,17 +69,58 @@ class TestAssertSwaDivergenceObserved(CustomTestCase):
max_retries=1,
)
def test_assert_swa_divergence_observed_uses_latest_line(self) -> None:
log = _GOOD_LINE + "\n" + _LATER_LINE + "\n"
harness, patcher = self._make_harness(log)
def test_assert_swa_divergence_observed_uses_peak_out_of_window(self) -> None:
diverged = SwaDivergenceLog(
forward_ct=120,
verify_full=10000,
verify_swa=4200,
swa_full_idx_divergence=512,
swa_out_of_window_tokens=8192,
).format()
trailing_zero = SwaDivergenceLog(
forward_ct=240,
verify_full=20000,
verify_swa=8400,
swa_full_idx_divergence=0,
swa_out_of_window_tokens=0,
).format()
harness, patcher = self._make_harness(diverged + "\n" + trailing_zero + "\n")
with patcher:
harness.assert_swa_divergence_observed(
min_swa_full_idx_divergence=1000,
min_swa_out_of_window_tokens=1,
min_swa_full_idx_divergence=1,
require_verify_lag=True,
flush_wait_seconds=0.0,
max_retries=1,
)
def test_assert_swa_divergence_observed_checks_verify_lag_on_latest_line(
self,
) -> None:
lagging = SwaDivergenceLog(
forward_ct=120,
verify_full=10000,
verify_swa=4200,
swa_full_idx_divergence=512,
swa_out_of_window_tokens=8192,
).format()
no_lag = SwaDivergenceLog(
forward_ct=240,
verify_full=20000,
verify_swa=20000,
swa_full_idx_divergence=1024,
swa_out_of_window_tokens=16384,
).format()
harness, patcher = self._make_harness(lagging + "\n" + no_lag + "\n")
with patcher:
with self.assertRaisesRegex(AssertionError, "verify_swa=20000"):
harness.assert_swa_divergence_observed(
min_swa_full_idx_divergence=1,
require_verify_lag=True,
flush_wait_seconds=0.0,
max_retries=1,
)
def test_assert_swa_divergence_observed_raises_when_below_threshold(self) -> None:
zero_mapping_line = SwaDivergenceLog(
forward_ct=100,
@@ -410,6 +410,40 @@ class TestSwaDivergenceReporterWithCompute(CustomTestCase):
self.assertEqual(parsed.verify_swa, 3)
class TestSwaDivergenceLogFindAll(CustomTestCase):
def test_find_all_returns_every_sample_in_order(self) -> None:
text = "\n".join(
SwaDivergenceLog(
forward_ct=ct,
verify_full=100 * ct,
verify_swa=10 * ct,
swa_full_idx_divergence=ct,
swa_out_of_window_tokens=0,
).format()
for ct in (20, 40, 60)
)
parsed = SwaDivergenceLog.find_all(text)
self.assertEqual([p.forward_ct for p, _ in parsed], [20, 40, 60])
def test_find_all_peak_survives_trailing_zero_sample(self) -> None:
text = "\n".join(
SwaDivergenceLog(
forward_ct=ct,
verify_full=1,
verify_swa=0,
swa_full_idx_divergence=1,
swa_out_of_window_tokens=oow,
).format()
for ct, oow in ((20, 0), (40, 4080), (60, 0))
)
parsed = SwaDivergenceLog.find_all(text)
self.assertEqual(max(p.swa_out_of_window_tokens for p, _ in parsed), 4080)
self.assertEqual(parsed[-1][0].swa_out_of_window_tokens, 0)
def test_find_all_returns_empty_list_when_no_lines(self) -> None:
self.assertEqual(SwaDivergenceLog.find_all("nothing here\n"), [])
class TestCanaryManagerSwaDivergenceWiring(CanaryManagerTestCase):
def test_swa_divergence_report_is_none_when_env_disabled(self) -> None:
with envs.SGLANG_KV_CANARY_SWA_DIVERGENCE_STATS_INTERVAL.override(