[UnifiedRadixTree]: Add KL accuracy CI for UnifiedTree with HiCache (#24346)
This commit is contained in:
@@ -1244,12 +1244,6 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
self.dec_lock_ref(ancestor_node)
|
self.dec_lock_ref(ancestor_node)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"load_back: kv_tokens=%d, node_id=%d",
|
|
||||||
kv_tokens,
|
|
||||||
last_hit_node.id,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Load H→D
|
# Load H→D
|
||||||
aux_xfers = [x for xfers in comp_xfers.values() for x in xfers]
|
aux_xfers = [x for xfers in comp_xfers.values() for x in xfers]
|
||||||
aux_xfers.extend(anchor_kv_shared_indices_xfers)
|
aux_xfers.extend(anchor_kv_shared_indices_xfers)
|
||||||
@@ -1367,14 +1361,9 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
req = params.req
|
req = params.req
|
||||||
|
|
||||||
if last_node.evicted or params.host_hit_length > 0:
|
if last_node.evicted or params.host_hit_length > 0:
|
||||||
logger.info(
|
|
||||||
"init_load_back triggered: node_id=%d, host_hit_length=%d",
|
|
||||||
last_node.id,
|
|
||||||
params.host_hit_length,
|
|
||||||
)
|
|
||||||
loading_values = self.load_back(last_node, mem_quota, req=req)
|
loading_values = self.load_back(last_node, mem_quota, req=req)
|
||||||
if loading_values is not None:
|
if loading_values is not None:
|
||||||
logger.info(
|
logger.debug(
|
||||||
"init_load_back success: loaded %d tokens for node %d",
|
"init_load_back success: loaded %d tokens for node %d",
|
||||||
len(loading_values),
|
len(loading_values),
|
||||||
last_node.id,
|
last_node.id,
|
||||||
|
|||||||
@@ -118,8 +118,13 @@ def compare_kl_divergence(
|
|||||||
|
|
||||||
|
|
||||||
# Common request helpers
|
# Common request helpers
|
||||||
def _flush_cache(base_url):
|
def _flush_cache(base_url, timeout_s=30):
|
||||||
requests.post(base_url + "/flush_cache")
|
response = requests.post(
|
||||||
|
base_url + "/flush_cache",
|
||||||
|
params={"timeout": timeout_s},
|
||||||
|
timeout=timeout_s + 10,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
|
||||||
|
|
||||||
def _generate(
|
def _generate(
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ MAMBA_TRACK_INTERVAL = 128
|
|||||||
SWA_MODEL = "openai/gpt-oss-20b"
|
SWA_MODEL = "openai/gpt-oss-20b"
|
||||||
FULL_MODEL = "Qwen/Qwen3-32B"
|
FULL_MODEL = "Qwen/Qwen3-32B"
|
||||||
|
|
||||||
register_cuda_ci(est_time=632, suite="stage-c-test-4-gpu-h100")
|
register_cuda_ci(est_time=760, suite="stage-c-test-4-gpu-h100")
|
||||||
|
|
||||||
|
|
||||||
class UnifiedRadixTreeTestMixin:
|
class UnifiedRadixTreeTestMixin:
|
||||||
@@ -235,8 +235,8 @@ class TestUnifiedSWARadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
|||||||
"""SWA hybrid + UnifiedRadixCache."""
|
"""SWA hybrid + UnifiedRadixCache."""
|
||||||
|
|
||||||
kl_threshold = 0.03
|
kl_threshold = 0.03
|
||||||
gsm8k_threshold = 0.75
|
gsm8k_threshold = 0.7
|
||||||
mmlu_threshold = 0.75
|
mmlu_threshold = 0.7
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
@@ -253,7 +253,7 @@ class TestUnifiedSWARadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
|||||||
"0.7",
|
"0.7",
|
||||||
"--disable-piecewise-cuda-graph",
|
"--disable-piecewise-cuda-graph",
|
||||||
],
|
],
|
||||||
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "0"},
|
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
|
||||||
)
|
)
|
||||||
cls.input_ids = get_input_ids(cls.model, num_samples=18)
|
cls.input_ids = get_input_ids(cls.model, num_samples=18)
|
||||||
|
|
||||||
@@ -262,54 +262,57 @@ class TestUnifiedSWARadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
|||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
# TODO(hzh): Currently, HiCache is not fully compatible with the CI CUDA13 environment; we need to wait for the fix before re-enabling the tests below.
|
class TestUnifiedMambaRadixCacheWithHiCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
||||||
# class TestUnifiedMambaRadixCacheWithHiCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
"""Mamba hybrid + HiCache + UnifiedRadixCache."""
|
||||||
# """Mamba hybrid + UnifiedRadixCache."""
|
|
||||||
#
|
kl_threshold = 0.003
|
||||||
# kl_threshold = 0.003
|
prefill_cache_assert = staticmethod(
|
||||||
# prefill_cache_assert = staticmethod(
|
make_mamba_prefill_assert(chunk_size=MAMBA_CHUNK_SIZE)
|
||||||
# make_mamba_prefill_assert(chunk_size=MAMBA_CHUNK_SIZE)
|
)
|
||||||
# )
|
decode_cache_assert = staticmethod(
|
||||||
# decode_cache_assert = staticmethod(
|
make_mamba_decode_assert(track_interval=MAMBA_TRACK_INTERVAL)
|
||||||
# make_mamba_decode_assert(track_interval=MAMBA_TRACK_INTERVAL)
|
)
|
||||||
# )
|
|
||||||
#
|
@classmethod
|
||||||
# @classmethod
|
def setUpClass(cls):
|
||||||
# def setUpClass(cls):
|
cls.model = MAMBA_MODEL
|
||||||
# cls.model = MAMBA_MODEL
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
# cls.base_url = DEFAULT_URL_FOR_TEST
|
cls.process = popen_launch_server(
|
||||||
# cls.process = popen_launch_server(
|
cls.model,
|
||||||
# cls.model,
|
cls.base_url,
|
||||||
# cls.base_url,
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
# timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
other_args=[
|
||||||
# other_args=[
|
"--tp-size",
|
||||||
# "--tp-size",
|
"4",
|
||||||
# "4",
|
"--chunked-prefill-size",
|
||||||
# "--chunked-prefill-size",
|
"2048",
|
||||||
# "2048",
|
"--mem-fraction-static",
|
||||||
# "--mem-fraction-static",
|
"0.85",
|
||||||
# "0.85",
|
"--mamba-scheduler-strategy",
|
||||||
# "--mamba-scheduler-strategy",
|
"extra_buffer",
|
||||||
# "extra_buffer",
|
"--mamba-track-interval",
|
||||||
# "--mamba-track-interval",
|
str(MAMBA_TRACK_INTERVAL),
|
||||||
# str(MAMBA_TRACK_INTERVAL),
|
"--enable-hierarchical-cache",
|
||||||
# "--enable-hierarchical-cache",
|
"--hicache-ratio",
|
||||||
# "--hicache-ratio",
|
"4",
|
||||||
# "1.5",
|
"--hicache-write-policy",
|
||||||
# "--hicache-write-policy",
|
"write_through",
|
||||||
# "write_through",
|
"--hicache-io-backend",
|
||||||
# "--hicache-io-backend",
|
"direct",
|
||||||
# "direct",
|
"--hicache-mem-layout",
|
||||||
# "--hicache-mem-layout",
|
"page_first_direct",
|
||||||
# "page_first_direct",
|
"--max-total-tokens",
|
||||||
# ],
|
"12000",
|
||||||
# env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
|
"--max-running-requests",
|
||||||
# )
|
"4",
|
||||||
# cls.input_ids = get_input_ids(cls.model, num_samples=18)
|
],
|
||||||
#
|
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
|
||||||
# @classmethod
|
)
|
||||||
# def tearDownClass(cls):
|
cls.input_ids = get_input_ids(cls.model, num_samples=18)
|
||||||
# kill_process_tree(cls.process.pid)
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user