Add bit-exact hicache logprob-consistency test (#34356)
This commit is contained in:
@@ -12,6 +12,7 @@ cards.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
import random
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
@@ -19,7 +20,18 @@ from urllib.parse import urlparse
|
|||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
# Aliased for the same reason as the single-turn helpers above.
|
||||||
|
from sglang.test.kl_multiturn_utils import (
|
||||||
|
make_mamba_decode_assert,
|
||||||
|
)
|
||||||
|
from sglang.test.kl_multiturn_utils import (
|
||||||
|
test_input_output_logprobs_match_decode_cache_hit_helper as assert_multiturn_decode_cache_hit,
|
||||||
|
)
|
||||||
|
|
||||||
# Aliased so pytest does not collect the imported `test_`-prefixed helpers as tests.
|
# Aliased so pytest does not collect the imported `test_`-prefixed helpers as tests.
|
||||||
|
from sglang.test.kl_test_utils import (
|
||||||
|
get_input_ids,
|
||||||
|
)
|
||||||
from sglang.test.kl_test_utils import (
|
from sglang.test.kl_test_utils import (
|
||||||
test_input_output_logprobs_match_decode_cache_hit_helper as assert_logprobs_match_decode_cache_hit,
|
test_input_output_logprobs_match_decode_cache_hit_helper as assert_logprobs_match_decode_cache_hit,
|
||||||
)
|
)
|
||||||
@@ -36,7 +48,7 @@ from sglang.test.test_utils import (
|
|||||||
popen_launch_server,
|
popen_launch_server,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=1200, stage="extra-b", runner_config="4-gpu-b200")
|
register_cuda_ci(est_time=1700, stage="extra-b", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
_MODEL_PATH = os.environ.get(
|
_MODEL_PATH = os.environ.get(
|
||||||
"INKLING_SMALL_TEST_MODEL_PATH", "thinkingmachines/Inkling-Small-NVFP4"
|
"INKLING_SMALL_TEST_MODEL_PATH", "thinkingmachines/Inkling-Small-NVFP4"
|
||||||
@@ -203,5 +215,123 @@ class TestInklingSmallNvfp4Deterministic(CustomTestCase):
|
|||||||
self._run(assert_logprobs_match_decode_cache_hit, min_cache_hit_ratio=0.99)
|
self._run(assert_logprobs_match_decode_cache_hit, min_cache_hit_ratio=0.99)
|
||||||
|
|
||||||
|
|
||||||
|
# The multi-turn branching harness, unlike the single-turn helpers above, replays
|
||||||
|
# turn N on top of turn N-1's history, so cache hits land at many different
|
||||||
|
# prefix lengths and nine interleaved branches share prefixes before diverging.
|
||||||
|
# That is what reaches the decode-side track save.
|
||||||
|
KL_HICACHE_TRACK_INTERVAL = 128
|
||||||
|
|
||||||
|
|
||||||
|
def _random_suffixes(n: int, length: int, seed: int) -> list[list[int]]:
|
||||||
|
rng = random.Random(seed)
|
||||||
|
return [[rng.randint(1, 30000) for _ in range(length)] for _ in range(n)]
|
||||||
|
|
||||||
|
|
||||||
|
class TestInklingSmallNvfp4HiCacheDeterministic(CustomTestCase):
|
||||||
|
"""HiCache round trip must be bit exact, since nothing about moving a state
|
||||||
|
to host memory and back is allowed to change it.
|
||||||
|
|
||||||
|
Every knob here is load bearing, and dropping any one of them takes the
|
||||||
|
measurement back to zero even when the state layer is wrong:
|
||||||
|
|
||||||
|
- hicache write_through, so an insert copies to host immediately
|
||||||
|
- the decode CUDA graph, since the corrupted slot is only read through the
|
||||||
|
captured decode path
|
||||||
|
- a short ``--mamba-track-interval``, so decode actually crosses a track
|
||||||
|
boundary within a turn
|
||||||
|
- pools tight enough to force eviction while requests are live
|
||||||
|
- multi-turn cache hits, because a wrong slot only surfaces once a later hit
|
||||||
|
restores it
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = _MODEL_PATH
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--tp",
|
||||||
|
"4",
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--quantization",
|
||||||
|
"modelopt_fp4",
|
||||||
|
"--attention-backend",
|
||||||
|
"fa4",
|
||||||
|
"--page-size",
|
||||||
|
"128",
|
||||||
|
"--fp4-gemm-backend",
|
||||||
|
"flashinfer_trtllm",
|
||||||
|
"--moe-runner-backend",
|
||||||
|
"flashinfer_trtllm_routed",
|
||||||
|
"--mamba-radix-cache-strategy",
|
||||||
|
"extra_buffer",
|
||||||
|
"--mamba-track-interval",
|
||||||
|
str(KL_HICACHE_TRACK_INTERVAL),
|
||||||
|
"--swa-full-tokens-ratio",
|
||||||
|
"0.1",
|
||||||
|
"--mamba-full-memory-ratio",
|
||||||
|
"0.1",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.85",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"2048",
|
||||||
|
"--enable-hierarchical-cache",
|
||||||
|
"--hicache-ratio",
|
||||||
|
"4",
|
||||||
|
"--hicache-write-policy",
|
||||||
|
"write_through",
|
||||||
|
"--hicache-io-backend",
|
||||||
|
"direct",
|
||||||
|
"--hicache-mem-layout",
|
||||||
|
"page_first_direct",
|
||||||
|
"--max-total-tokens",
|
||||||
|
"65536",
|
||||||
|
"--max-mamba-cache-size",
|
||||||
|
"500",
|
||||||
|
"--max-running-requests",
|
||||||
|
"4",
|
||||||
|
"--enable-deterministic-inference",
|
||||||
|
],
|
||||||
|
env={**os.environ, "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
|
||||||
|
)
|
||||||
|
cls.input_ids = get_input_ids(cls.model, num_samples=18, trust_remote_code=True)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if getattr(cls, "process", None) is not None:
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def test_multiturn_decode_cache_hit_over_hicache(self):
|
||||||
|
"""Nine interleaved branches, three turns, decode hits served through the
|
||||||
|
host tier. Also asserts every hit lands on a track boundary, which the
|
||||||
|
logprob comparison cannot see on its own."""
|
||||||
|
groups, branches = 3, 3
|
||||||
|
n = groups * branches
|
||||||
|
first_turn = []
|
||||||
|
for g in range(groups):
|
||||||
|
base = self.input_ids[g][:512]
|
||||||
|
for _ in range(branches):
|
||||||
|
first_turn.append(list(base))
|
||||||
|
|
||||||
|
assert_multiturn_decode_cache_hit(
|
||||||
|
self.base_url,
|
||||||
|
self.model,
|
||||||
|
KL_DIV_THRESHOLD,
|
||||||
|
first_turn,
|
||||||
|
turn_suffixes=[
|
||||||
|
_random_suffixes(n, 512, seed=300),
|
||||||
|
_random_suffixes(n, 256, seed=400),
|
||||||
|
],
|
||||||
|
assert_decode_cached_tokens=make_mamba_decode_assert(
|
||||||
|
track_interval=KL_HICACHE_TRACK_INTERVAL
|
||||||
|
),
|
||||||
|
branches_per_group=branches,
|
||||||
|
max_new_tokens=KL_MAX_NEW_TOKENS,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user