[UnifiedTree]: Sync mamba int8 checkpoint (#30626)

This commit is contained in:
Zhangheng
2026-07-10 23:46:38 +08:00
committed by GitHub
parent 559854fe6a
commit 0299393758
3 changed files with 181 additions and 24 deletions
@@ -174,7 +174,7 @@ class MambaComponent(TreeComponent):
# Device layer
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)
self.cache.component_evictable_size_[self.component_type] -= freed
cd.value = None
@@ -294,6 +294,33 @@ class MambaComponent(TreeComponent):
assert slot is not None, "Can not alloc mamba cache"
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(
self,
req: Req,
@@ -326,18 +353,35 @@ class MambaComponent(TreeComponent):
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx(
req
)
mamba_value = (
active_value = (
req.mamba_ping_pong_track_buffer[keep_idx].unsqueeze(-1).clone()
)
else:
mamba_value = req.mamba_pool_idx.unsqueeze(-1).clone()
insert_params.mamba_value = mamba_value
active_value = req.mamba_pool_idx.unsqueeze(-1).clone()
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
else:
if cache_len is None:
return 0
# 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()
mamba_value_donated = (
self.cache.req_to_token_pool.donate_mamba_ping_pong_slot(
@@ -364,29 +408,40 @@ class MambaComponent(TreeComponent):
insert_params: Optional[InsertParams] = None,
) -> None:
if is_finished:
mamba_exist = (
insert_result.mamba_exist if insert_result is not None else True
mamba_value_inserted = (
insert_result is not None and not insert_result.mamba_exist
)
if self.enable_mamba_extra_buffer:
keep_idx = self.cache.req_to_token_pool.get_mamba_ping_pong_keep_idx(
req
pool = self.cache.req_to_token_pool
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:
keep_idx = None
if mamba_exist:
keep_idx = None
free_mamba_cache = True if self.enable_mamba_extra_buffer else mamba_exist
if free_mamba_cache:
self.cache.req_to_token_pool.free_mamba_cache(
if insert_value_unused:
self._free_mamba_value(insert_params.mamba_value)
pool.free_mamba_cache(req)
return
if self.enable_mamba_extra_buffer:
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
)
return
if not mamba_value_inserted:
pool.free_mamba_cache(req)
else:
if insert_params.mamba_value is not None and (
insert_result is None or insert_result.mamba_exist
):
self.cache.req_to_token_pool.mamba_allocator.free(
insert_params.mamba_value
)
self._free_mamba_value(insert_params.mamba_value)
req.mamba_last_track_seqlen = None
# ---- HiCache Hooks ----