[UnifiedTree]: Sync mamba int8 checkpoint (#30626)
This commit is contained in:
@@ -174,7 +174,7 @@ class MambaComponent(TreeComponent):
|
|||||||
|
|
||||||
# Device layer
|
# Device layer
|
||||||
if EvictLayer.DEVICE in target and cd.value is not None:
|
if EvictLayer.DEVICE in target and cd.value is not None:
|
||||||
self.cache.req_to_token_pool.mamba_allocator.free(cd.value)
|
self._free_mamba_value(cd.value)
|
||||||
freed = len(cd.value)
|
freed = len(cd.value)
|
||||||
self.cache.component_evictable_size_[self.component_type] -= freed
|
self.cache.component_evictable_size_[self.component_type] -= freed
|
||||||
cd.value = None
|
cd.value = None
|
||||||
@@ -294,6 +294,33 @@ class MambaComponent(TreeComponent):
|
|||||||
assert slot is not None, "Can not alloc mamba cache"
|
assert slot is not None, "Can not alloc mamba cache"
|
||||||
return slot
|
return slot
|
||||||
|
|
||||||
|
@property
|
||||||
|
def int8_ckpt_pool(self):
|
||||||
|
return getattr(self.cache.req_to_token_pool, "mamba_ckpt_pool", None)
|
||||||
|
|
||||||
|
def _alloc_int8_ckpt_slot(self) -> torch.Tensor:
|
||||||
|
slot = self.int8_ckpt_pool.alloc(1)
|
||||||
|
if slot is None:
|
||||||
|
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
|
||||||
|
slot = self.int8_ckpt_pool.alloc(1)
|
||||||
|
assert slot is not None, "Can not alloc int8 mamba checkpoint slot"
|
||||||
|
return slot
|
||||||
|
|
||||||
|
def _commit_int8_checkpoint(self, active_slots: torch.Tensor) -> torch.Tensor:
|
||||||
|
ckpt_slot = self._alloc_int8_ckpt_slot()
|
||||||
|
self.int8_ckpt_pool.store_from_active(
|
||||||
|
self.cache.req_to_token_pool.mamba_pool,
|
||||||
|
active_slots.view(-1),
|
||||||
|
ckpt_slot,
|
||||||
|
)
|
||||||
|
return ckpt_slot
|
||||||
|
|
||||||
|
def _free_mamba_value(self, mamba_value: torch.Tensor) -> None:
|
||||||
|
if self.int8_ckpt_pool is not None:
|
||||||
|
self.int8_ckpt_pool.free(mamba_value)
|
||||||
|
else:
|
||||||
|
self.cache.req_to_token_pool.mamba_allocator.free(mamba_value)
|
||||||
|
|
||||||
def prepare_for_caching_req(
|
def prepare_for_caching_req(
|
||||||
self,
|
self,
|
||||||
req: Req,
|
req: Req,
|
||||||
@@ -326,18 +353,35 @@ class MambaComponent(TreeComponent):
|
|||||||
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx(
|
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx(
|
||||||
req
|
req
|
||||||
)
|
)
|
||||||
mamba_value = (
|
active_value = (
|
||||||
req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
|
req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone()
|
active_value = req.mamba_pool_idx.unsqueeze(-1).clone()
|
||||||
insert_params.mamba_value = mamba_value
|
if self.int8_ckpt_pool is not None:
|
||||||
|
insert_params.mamba_value = self._commit_int8_checkpoint(active_value)
|
||||||
|
else:
|
||||||
|
insert_params.mamba_value = active_value
|
||||||
return cache_len
|
return cache_len
|
||||||
else:
|
else:
|
||||||
if cache_len is None:
|
if cache_len is None:
|
||||||
return 0
|
return 0
|
||||||
# Donate the mamba index to the radix cache instead of copying.
|
# Donate the mamba index to the radix cache instead of copying.
|
||||||
if self.enable_mamba_extra_buffer:
|
if self.int8_ckpt_pool is not None:
|
||||||
|
if self.enable_mamba_extra_buffer:
|
||||||
|
new_slot = self._alloc_mamba_slot()
|
||||||
|
src_active = (
|
||||||
|
self.cache.req_to_token_pool.donate_mamba_ping_pong_slot(
|
||||||
|
req, new_slot
|
||||||
|
)
|
||||||
|
)
|
||||||
|
mamba_value_donated = self._commit_int8_checkpoint(src_active)
|
||||||
|
self.cache.req_to_token_pool.mamba_allocator.free(src_active)
|
||||||
|
else:
|
||||||
|
mamba_value_donated = self._commit_int8_checkpoint(
|
||||||
|
req.mamba_pool_idx.view(-1)
|
||||||
|
)
|
||||||
|
elif self.enable_mamba_extra_buffer:
|
||||||
new_slot = self._alloc_mamba_slot()
|
new_slot = self._alloc_mamba_slot()
|
||||||
mamba_value_donated = (
|
mamba_value_donated = (
|
||||||
self.cache.req_to_token_pool.donate_mamba_ping_pong_slot(
|
self.cache.req_to_token_pool.donate_mamba_ping_pong_slot(
|
||||||
@@ -364,29 +408,40 @@ class MambaComponent(TreeComponent):
|
|||||||
insert_params: Optional[InsertParams] = None,
|
insert_params: Optional[InsertParams] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
if is_finished:
|
if is_finished:
|
||||||
mamba_exist = (
|
mamba_value_inserted = (
|
||||||
insert_result.mamba_exist if insert_result is not None else True
|
insert_result is not None and not insert_result.mamba_exist
|
||||||
)
|
)
|
||||||
if self.enable_mamba_extra_buffer:
|
pool = self.cache.req_to_token_pool
|
||||||
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx(
|
|
||||||
req
|
if self.int8_ckpt_pool is not None:
|
||||||
|
insert_value_unused = (
|
||||||
|
not mamba_value_inserted
|
||||||
|
and insert_params is not None
|
||||||
|
and insert_params.mamba_value is not None
|
||||||
)
|
)
|
||||||
else:
|
if insert_value_unused:
|
||||||
keep_idx = None
|
self._free_mamba_value(insert_params.mamba_value)
|
||||||
if mamba_exist:
|
pool.free_mamba_cache(req)
|
||||||
keep_idx = None
|
return
|
||||||
free_mamba_cache = True if self.enable_mamba_extra_buffer else mamba_exist
|
|
||||||
if free_mamba_cache:
|
if self.enable_mamba_extra_buffer:
|
||||||
self.cache.req_to_token_pool.free_mamba_cache(
|
keep_idx = (
|
||||||
|
pool.get_mamba_ping_pong_keep_idx(req)
|
||||||
|
if mamba_value_inserted
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
pool.free_mamba_cache(
|
||||||
req, mamba_ping_pong_track_buffer_to_keep=keep_idx
|
req, mamba_ping_pong_track_buffer_to_keep=keep_idx
|
||||||
)
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if not mamba_value_inserted:
|
||||||
|
pool.free_mamba_cache(req)
|
||||||
else:
|
else:
|
||||||
if insert_params.mamba_value is not None and (
|
if insert_params.mamba_value is not None and (
|
||||||
insert_result is None or insert_result.mamba_exist
|
insert_result is None or insert_result.mamba_exist
|
||||||
):
|
):
|
||||||
self.cache.req_to_token_pool.mamba_allocator.free(
|
self._free_mamba_value(insert_params.mamba_value)
|
||||||
insert_params.mamba_value
|
|
||||||
)
|
|
||||||
req.mamba_last_track_seqlen = None
|
req.mamba_last_track_seqlen = None
|
||||||
|
|
||||||
# ---- HiCache Hooks ----
|
# ---- HiCache Hooks ----
|
||||||
|
|||||||
@@ -21,16 +21,24 @@ Usage:
|
|||||||
python3 -m unittest test_int8_mamba_checkpoint_e2e
|
python3 -m unittest test_int8_mamba_checkpoint_e2e
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
|
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
|
||||||
from sglang.test.kits.kl_divergence_kit import KLDivergenceMixin
|
from sglang.test.kits.kl_divergence_kit import KLDivergenceMixin
|
||||||
from sglang.test.server_fixtures.default_fixture import DefaultServerBase
|
from sglang.test.server_fixtures.default_fixture import (
|
||||||
from sglang.test.test_utils import DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST
|
DefaultServerBase,
|
||||||
|
openai_api_env,
|
||||||
|
)
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_HYBRID_MAMBA_MODEL_NAME_FOR_TEST,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=400, stage="extra-b", runner_config="4-gpu-h100")
|
register_cuda_ci(est_time=800, stage="extra-b", runner_config="4-gpu-h100")
|
||||||
|
|
||||||
|
|
||||||
class TestInt8MambaCheckpointE2E(KLDivergenceMixin, DefaultServerBase):
|
class TestInt8MambaCheckpointE2E(KLDivergenceMixin, DefaultServerBase):
|
||||||
@@ -89,5 +97,27 @@ class TestInt8MambaCheckpointE2E(KLDivergenceMixin, DefaultServerBase):
|
|||||||
self.assertGreaterEqual(metrics["accuracy"], self.gsm8k_threshold)
|
self.assertGreaterEqual(metrics["accuracy"], self.gsm8k_threshold)
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedRadixTreeInt8MambaCheckpointE2E(TestInt8MambaCheckpointE2E):
|
||||||
|
"""Run the same int8 mamba checkpoint checks with UnifiedRadixTree forced on."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
assert cls.model is not None, "Please set cls.model in subclass"
|
||||||
|
|
||||||
|
with openai_api_env(cls.api_key):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=cls.timeout,
|
||||||
|
other_args=cls.other_args,
|
||||||
|
env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"},
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid, wait_timeout=60)
|
||||||
|
time.sleep(2)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -91,6 +91,7 @@ class CacheConfig:
|
|||||||
mamba_head_dim: int = 16
|
mamba_head_dim: int = 16
|
||||||
mamba_state_size: int = 16
|
mamba_state_size: int = 16
|
||||||
mamba_conv_kernel: int = 4
|
mamba_conv_kernel: int = 4
|
||||||
|
enable_int8_mamba_checkpoint: bool = False
|
||||||
|
|
||||||
# Model / pool
|
# Model / pool
|
||||||
kv_size: int = 256
|
kv_size: int = 256
|
||||||
@@ -214,7 +215,11 @@ class TestUnifiedTreeNodeGetPrefixHashValues(CustomTestCase):
|
|||||||
|
|
||||||
def build_fixture(cfg: CacheConfig, *, enable_kv_cache_events: bool = False):
|
def build_fixture(cfg: CacheConfig, *, enable_kv_cache_events: bool = False):
|
||||||
"""Create (tree, allocator, req_to_token_pool) from a CacheConfig."""
|
"""Create (tree, allocator, req_to_token_pool) from a CacheConfig."""
|
||||||
server_args = ServerArgs(model_path="dummy", page_size=cfg.page_size)
|
server_args = ServerArgs(
|
||||||
|
model_path="dummy",
|
||||||
|
page_size=cfg.page_size,
|
||||||
|
enable_int8_mamba_checkpoint=cfg.enable_int8_mamba_checkpoint,
|
||||||
|
)
|
||||||
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise
|
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise
|
||||||
# loads the HF config for self.model_path — impossible for the dummy model.
|
# loads the HF config for self.model_path — impossible for the dummy model.
|
||||||
# Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE.
|
# Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE.
|
||||||
@@ -4037,6 +4042,73 @@ class UnifiedLRUListBoundedRefreshTest(CustomTestCase):
|
|||||||
self.assertEqual(self._lru_order(lru), before)
|
self.assertEqual(self._lru_order(lru), before)
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedRadixCacheInt8MambaCheckpoint(CustomTestCase):
|
||||||
|
cfg = CacheConfig(
|
||||||
|
components=(ComponentType.FULL, ComponentType.MAMBA),
|
||||||
|
enable_mamba_extra_buffer=True,
|
||||||
|
enable_int8_mamba_checkpoint=True,
|
||||||
|
mamba_cache_size=8,
|
||||||
|
kv_size=64,
|
||||||
|
max_context_len=64,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _make_req(self, req_to_token_pool, tokens):
|
||||||
|
req = Req(
|
||||||
|
rid="int8-mamba",
|
||||||
|
origin_input_text="",
|
||||||
|
origin_input_ids=array("q", tokens),
|
||||||
|
sampling_params=SamplingParams(temperature=0, max_new_tokens=1),
|
||||||
|
)
|
||||||
|
req_to_token_pool.alloc([req])
|
||||||
|
req.output_ids = array("q")
|
||||||
|
req.kv_committed_len = len(tokens)
|
||||||
|
req.kv_allocated_len = len(tokens)
|
||||||
|
req.cache_protected_len = 0
|
||||||
|
req.swa_uuid_for_lock = None
|
||||||
|
req.extra_key = None
|
||||||
|
req.mamba_last_track_seqlen = len(tokens)
|
||||||
|
return req
|
||||||
|
|
||||||
|
def _cache_finished(self, cache, allocator, req_to_token_pool, tokens):
|
||||||
|
req = self._make_req(req_to_token_pool, tokens)
|
||||||
|
kv_indices = allocator.alloc(len(tokens))
|
||||||
|
self.assertIsNotNone(kv_indices)
|
||||||
|
req_to_token_pool.write((req.req_pool_idx, slice(0, len(tokens))), kv_indices)
|
||||||
|
req.last_node = cache.root_node
|
||||||
|
|
||||||
|
cache.cache_finished_req(req, is_insert=True)
|
||||||
|
|
||||||
|
def test_finished_req_stores_radix_mamba_state_in_int8_pool(self):
|
||||||
|
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
|
||||||
|
ckpt_pool = req_to_token_pool.mamba_ckpt_pool
|
||||||
|
self.assertIsNotNone(ckpt_pool)
|
||||||
|
active_initial = req_to_token_pool.mamba_allocator.available_size()
|
||||||
|
ckpt_initial = ckpt_pool.available_size()
|
||||||
|
tokens = [1, 2, 3, 4]
|
||||||
|
|
||||||
|
self._cache_finished(cache, allocator, req_to_token_pool, tokens)
|
||||||
|
self.assertEqual(
|
||||||
|
req_to_token_pool.mamba_allocator.available_size(), active_initial
|
||||||
|
)
|
||||||
|
self.assertEqual(ckpt_pool.available_size(), ckpt_initial - 1)
|
||||||
|
self.assertEqual(cache.mamba_evictable_size(), 1)
|
||||||
|
|
||||||
|
self._cache_finished(cache, allocator, req_to_token_pool, tokens)
|
||||||
|
self.assertEqual(
|
||||||
|
req_to_token_pool.mamba_allocator.available_size(), active_initial
|
||||||
|
)
|
||||||
|
self.assertEqual(ckpt_pool.available_size(), ckpt_initial - 1)
|
||||||
|
self.assertEqual(cache.mamba_evictable_size(), 1)
|
||||||
|
|
||||||
|
result = cache.evict(EvictParams(mamba_num=1))
|
||||||
|
self.assertEqual(result.mamba_num_evicted, 1)
|
||||||
|
self.assertEqual(
|
||||||
|
req_to_token_pool.mamba_allocator.available_size(), active_initial
|
||||||
|
)
|
||||||
|
self.assertEqual(ckpt_pool.available_size(), ckpt_initial)
|
||||||
|
self.assertEqual(cache.mamba_evictable_size(), 0)
|
||||||
|
|
||||||
|
|
||||||
_CONFIGS: list[CacheConfig] = [
|
_CONFIGS: list[CacheConfig] = [
|
||||||
CacheConfig(page_size=1, components=(ComponentType.FULL,)),
|
CacheConfig(page_size=1, components=(ComponentType.FULL,)),
|
||||||
CacheConfig(page_size=4, components=(ComponentType.FULL,)),
|
CacheConfig(page_size=4, components=(ComponentType.FULL,)),
|
||||||
|
|||||||
Reference in New Issue
Block a user