[HiCache] Fix Mamba track-boundary bookkeeping under overlap scheduling (#29792)
Co-authored-by: Hanming Lu <hanminglu@meta.com>
This commit is contained in:
co-authored by
Hanming Lu
parent
a3bd7d9401
commit
8f3d3a31f4
@@ -275,6 +275,7 @@ class MlxAuxiliaryStateReqToTokenPool(ReqToTokenPool):
|
|||||||
self.auxiliary_state_pool.free(track_buffer)
|
self.auxiliary_state_pool.free(track_buffer)
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.mamba_next_track_idx = None
|
||||||
|
req.mamba_last_track_idx = None
|
||||||
|
|
||||||
def free_auxiliary_state_cache(self, req, track_buffer_to_keep=None):
|
def free_auxiliary_state_cache(self, req, track_buffer_to_keep=None):
|
||||||
self.free_mamba_cache(
|
self.free_mamba_cache(
|
||||||
@@ -385,6 +386,7 @@ class MlxAuxiliaryStateComponent(MambaComponent):
|
|||||||
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
|
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.mamba_next_track_idx = None
|
||||||
|
req.mamba_last_track_idx = None
|
||||||
req.mamba_last_track_seqlen = None
|
req.mamba_last_track_seqlen = None
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -409,5 +411,6 @@ class MlxAuxiliaryStateComponent(MambaComponent):
|
|||||||
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
|
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.mamba_next_track_idx = None
|
||||||
|
req.mamba_last_track_idx = None
|
||||||
req.mamba_pool_idx = None
|
req.mamba_pool_idx = None
|
||||||
req.mamba_last_track_seqlen = None
|
req.mamba_last_track_seqlen = None
|
||||||
|
|||||||
@@ -399,6 +399,7 @@ class MlxModelRunner:
|
|||||||
return
|
return
|
||||||
req.mamba_ping_pong_track_buffer = track_buffer
|
req.mamba_ping_pong_track_buffer = track_buffer
|
||||||
req.mamba_next_track_idx = 0
|
req.mamba_next_track_idx = 0
|
||||||
|
req.mamba_last_track_idx = 0
|
||||||
|
|
||||||
pool.store_cache(
|
pool.store_cache(
|
||||||
track_buffer[0],
|
track_buffer[0],
|
||||||
|
|||||||
@@ -939,6 +939,7 @@ class Req(ReqDllmMixin):
|
|||||||
self.mamba_pool_idx: Optional[torch.Tensor] = None # shape (1)
|
self.mamba_pool_idx: Optional[torch.Tensor] = None # shape (1)
|
||||||
self.mamba_ping_pong_track_buffer: Optional[torch.Tensor] = None # shape (2)
|
self.mamba_ping_pong_track_buffer: Optional[torch.Tensor] = None # shape (2)
|
||||||
self.mamba_next_track_idx: Optional[int] = None # 0 or 1
|
self.mamba_next_track_idx: Optional[int] = None # 0 or 1
|
||||||
|
self.mamba_last_track_idx: Optional[int] = None # 0 or 1
|
||||||
self.mamba_last_track_seqlen: Optional[int] = (
|
self.mamba_last_track_seqlen: Optional[int] = (
|
||||||
None # seq len of the last cached mamba state
|
None # seq len of the last cached mamba state
|
||||||
)
|
)
|
||||||
@@ -1693,6 +1694,7 @@ class Req(ReqDllmMixin):
|
|||||||
self.mamba_pool_idx = None
|
self.mamba_pool_idx = None
|
||||||
self.mamba_ping_pong_track_buffer = None
|
self.mamba_ping_pong_track_buffer = None
|
||||||
self.mamba_next_track_idx = None
|
self.mamba_next_track_idx = None
|
||||||
|
self.mamba_last_track_idx = None
|
||||||
self.mamba_last_track_seqlen = None
|
self.mamba_last_track_seqlen = None
|
||||||
self.mamba_branching_seqlen = None
|
self.mamba_branching_seqlen = None
|
||||||
self.mamba_cow_src_index = None
|
self.mamba_cow_src_index = None
|
||||||
@@ -1883,6 +1885,7 @@ def set_mamba_track_indices_from_reqs(
|
|||||||
req.mamba_next_track_idx if req.mamba_next_track_idx is not None else 0
|
req.mamba_next_track_idx if req.mamba_next_track_idx is not None else 0
|
||||||
for req in batch.reqs
|
for req in batch.reqs
|
||||||
]
|
]
|
||||||
|
batch.mamba_track_buffer_indices = list(track_positions)
|
||||||
idx = (
|
idx = (
|
||||||
torch.tensor(
|
torch.tensor(
|
||||||
track_positions,
|
track_positions,
|
||||||
@@ -2079,6 +2082,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
|
|
||||||
# For hybrid GDN prefix cache
|
# For hybrid GDN prefix cache
|
||||||
mamba_track_indices: torch.Tensor = None # shape: [b], int64
|
mamba_track_indices: torch.Tensor = None # shape: [b], int64
|
||||||
|
# Per-batch snapshot of the logical ping-pong positions selected for this
|
||||||
|
# forward (normally req.mamba_next_track_idx; spec may override it). Result
|
||||||
|
# processing uses it to update req.mamba_last_track_idx, since both req-level
|
||||||
|
# indices may advance under overlap.
|
||||||
|
mamba_track_buffer_indices: Optional[List[int]] = None # shape: [b], 0 or 1
|
||||||
mamba_track_mask: torch.Tensor = None # shape: [b], bool
|
mamba_track_mask: torch.Tensor = None # shape: [b], bool
|
||||||
mamba_track_seqlens: torch.Tensor = None # shape: [b], int64
|
mamba_track_seqlens: torch.Tensor = None # shape: [b], int64
|
||||||
mamba_track_mask_cpu: Optional[List[bool]] = None # shape: [b]
|
mamba_track_mask_cpu: Optional[List[bool]] = None # shape: [b]
|
||||||
@@ -2666,6 +2674,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# In lazy mode, skip the swap — the second ping-pong slot is not
|
# In lazy mode, skip the swap — the second ping-pong slot is not
|
||||||
# allocated yet; it will be allocated on demand at the track boundary
|
# allocated yet; it will be allocated on demand at the track boundary
|
||||||
# in mamba_lazy_prealloc_at_boundary during prepare_for_decode.
|
# in mamba_lazy_prealloc_at_boundary during prepare_for_decode.
|
||||||
|
req.mamba_last_track_idx = req.mamba_next_track_idx
|
||||||
if not mamba_extra_buffer_lazy_enabled():
|
if not mamba_extra_buffer_lazy_enabled():
|
||||||
req.mamba_next_track_idx = (
|
req.mamba_next_track_idx = (
|
||||||
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
|
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
|
||||||
@@ -3067,6 +3076,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
self.mamba_track_indices = torch.empty(
|
self.mamba_track_indices = torch.empty(
|
||||||
(0,), dtype=torch.int64, device=self.device
|
(0,), dtype=torch.int64, device=self.device
|
||||||
)
|
)
|
||||||
|
self.mamba_track_buffer_indices = []
|
||||||
else:
|
else:
|
||||||
if mamba_extra_buffer_lazy_enabled():
|
if mamba_extra_buffer_lazy_enabled():
|
||||||
self.mamba_lazy_prealloc_at_boundary(mamba_track_interval)
|
self.mamba_lazy_prealloc_at_boundary(mamba_track_interval)
|
||||||
@@ -3148,6 +3158,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
self.seq_lens_cpu = self.seq_lens_cpu[keep_indices]
|
self.seq_lens_cpu = self.seq_lens_cpu[keep_indices]
|
||||||
|
|
||||||
self.mamba_track_indices = None
|
self.mamba_track_indices = None
|
||||||
|
self.mamba_track_buffer_indices = None
|
||||||
self.mamba_track_mask = None
|
self.mamba_track_mask = None
|
||||||
self.mamba_track_seqlens = None
|
self.mamba_track_seqlens = None
|
||||||
self.mamba_track_mask_cpu = None
|
self.mamba_track_mask_cpu = None
|
||||||
@@ -3211,6 +3222,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
else:
|
else:
|
||||||
self.seq_lens_cpu = torch.cat([self.seq_lens_cpu, other.seq_lens_cpu])
|
self.seq_lens_cpu = torch.cat([self.seq_lens_cpu, other.seq_lens_cpu])
|
||||||
self.mamba_track_indices = None
|
self.mamba_track_indices = None
|
||||||
|
self.mamba_track_buffer_indices = None
|
||||||
self.mamba_track_mask = None
|
self.mamba_track_mask = None
|
||||||
self.mamba_track_seqlens = None
|
self.mamba_track_seqlens = None
|
||||||
self.mamba_track_mask_cpu = None
|
self.mamba_track_mask_cpu = None
|
||||||
@@ -3272,6 +3284,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
seq_lens_cpu=self.seq_lens_cpu,
|
seq_lens_cpu=self.seq_lens_cpu,
|
||||||
enable_overlap=self.enable_overlap,
|
enable_overlap=self.enable_overlap,
|
||||||
mamba_track_indices=self.mamba_track_indices,
|
mamba_track_indices=self.mamba_track_indices,
|
||||||
|
mamba_track_buffer_indices=self.mamba_track_buffer_indices,
|
||||||
mamba_track_mask=self.mamba_track_mask,
|
mamba_track_mask=self.mamba_track_mask,
|
||||||
mamba_track_seqlens=self.mamba_track_seqlens,
|
mamba_track_seqlens=self.mamba_track_seqlens,
|
||||||
mamba_track_mask_cpu=self.mamba_track_mask_cpu,
|
mamba_track_mask_cpu=self.mamba_track_mask_cpu,
|
||||||
|
|||||||
@@ -1016,8 +1016,12 @@ class SchedulerBatchResultProcessor:
|
|||||||
i: int,
|
i: int,
|
||||||
logits_output: LogitsProcessorOutput,
|
logits_output: LogitsProcessorOutput,
|
||||||
):
|
):
|
||||||
|
lazy = mamba_extra_buffer_lazy_enabled()
|
||||||
known_mamba_boundary = None
|
known_mamba_boundary = None
|
||||||
|
completed_mamba_boundary = None
|
||||||
|
lookahead = 0
|
||||||
if batch.mamba_track_mask_cpu is not None:
|
if batch.mamba_track_mask_cpu is not None:
|
||||||
|
completed_mamba_boundary = bool(batch.mamba_track_mask_cpu[i])
|
||||||
lookahead = req.decode_batch_idx - batch.mamba_decode_batch_idx_cpu[i]
|
lookahead = req.decode_batch_idx - batch.mamba_decode_batch_idx_cpu[i]
|
||||||
assert lookahead in (0, 1), (
|
assert lookahead in (0, 1), (
|
||||||
f"mamba result lookahead={lookahead} for req {req.rid}; "
|
f"mamba result lookahead={lookahead} for req {req.rid}; "
|
||||||
@@ -1028,15 +1032,28 @@ class SchedulerBatchResultProcessor:
|
|||||||
else:
|
else:
|
||||||
known_mamba_boundary = bool(batch.mamba_track_mask_next_cpu[i])
|
known_mamba_boundary = bool(batch.mamba_track_mask_next_cpu[i])
|
||||||
|
|
||||||
|
if completed_mamba_boundary and not lazy:
|
||||||
|
req.mamba_last_track_idx = batch.mamba_track_buffer_indices[i]
|
||||||
|
req.mamba_last_track_seqlen = req.kv_committed_len - lookahead
|
||||||
|
elif (
|
||||||
|
req.finished()
|
||||||
|
and lazy
|
||||||
|
and lookahead == 1
|
||||||
|
and known_mamba_boundary
|
||||||
|
and req.mamba_next_track_idx == req.mamba_last_track_idx
|
||||||
|
):
|
||||||
|
req.mamba_lazy_is_insert = False
|
||||||
|
|
||||||
# Called here (after update_finish_state) so req.finished() is valid
|
# Called here (after update_finish_state) so req.finished() is valid
|
||||||
# for mamba_lazy_post_decode_at_boundary inside.
|
# for mamba_lazy_post_decode_at_boundary inside.
|
||||||
if known_mamba_boundary is None or known_mamba_boundary:
|
should_update = completed_mamba_boundary if lazy else known_mamba_boundary
|
||||||
|
if should_update is None or should_update:
|
||||||
self._mamba_prefix_cache_update(
|
self._mamba_prefix_cache_update(
|
||||||
req,
|
req,
|
||||||
batch,
|
batch,
|
||||||
result,
|
result,
|
||||||
i,
|
i,
|
||||||
known_boundary=known_mamba_boundary is True,
|
known_boundary=not lazy and known_mamba_boundary is True,
|
||||||
)
|
)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
@@ -1129,14 +1146,18 @@ class SchedulerBatchResultProcessor:
|
|||||||
if not at_boundary:
|
if not at_boundary:
|
||||||
return
|
return
|
||||||
|
|
||||||
req.mamba_last_track_seqlen = track_seqlen
|
track_idx = req.mamba_next_track_idx
|
||||||
|
if not known_boundary and batch.mamba_track_buffer_indices is not None:
|
||||||
|
track_idx = batch.mamba_track_buffer_indices[i]
|
||||||
|
if not known_boundary:
|
||||||
|
req.mamba_last_track_seqlen = track_seqlen
|
||||||
if lazy:
|
if lazy:
|
||||||
self.mamba_lazy_post_decode_at_boundary(req, batch)
|
self.mamba_lazy_post_decode_at_boundary(req, batch, track_idx)
|
||||||
else:
|
else:
|
||||||
|
if not known_boundary:
|
||||||
|
req.mamba_last_track_idx = track_idx
|
||||||
req.mamba_next_track_idx = (
|
req.mamba_next_track_idx = (
|
||||||
batch.req_to_token_pool.get_mamba_ping_pong_other_idx(
|
batch.req_to_token_pool.get_mamba_ping_pong_other_idx(track_idx)
|
||||||
req.mamba_next_track_idx
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _mamba_lazy_spec_update(
|
def _mamba_lazy_spec_update(
|
||||||
@@ -1193,6 +1214,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
req.mamba_next_track_idx = planned_pos
|
req.mamba_next_track_idx = planned_pos
|
||||||
# else: in-place fallback, or promoted by an earlier confirmation —
|
# else: in-place fallback, or promoted by an earlier confirmation —
|
||||||
# keep holds the track_seqlen state either way.
|
# keep holds the track_seqlen state either way.
|
||||||
|
req.mamba_last_track_idx = planned_pos
|
||||||
req.mamba_last_track_seqlen = track_seqlen
|
req.mamba_last_track_seqlen = track_seqlen
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -1216,9 +1238,8 @@ class SchedulerBatchResultProcessor:
|
|||||||
matches what the forward's tracking mask used:
|
matches what the forward's tracking mask used:
|
||||||
``prepare_for_decode`` increments both ``seq_lens_cpu`` and
|
``prepare_for_decode`` increments both ``seq_lens_cpu`` and
|
||||||
``kv_committed_len`` by 1, then checks
|
``kv_committed_len`` by 1, then checks
|
||||||
``seq_lens_cpu % interval == 0``. Using ``kv_committed_len``
|
``seq_lens_cpu % interval == 0``. Subtracting the overlap
|
||||||
here reproduces that check exactly, and the value is always a
|
lookahead from ``kv_committed_len`` reproduces that check.
|
||||||
multiple of ``interval`` (hence page-aligned).
|
|
||||||
|
|
||||||
For spec decode, the boundary is detected by comparing the
|
For spec decode, the boundary is detected by comparing the
|
||||||
accepted seq_len range against interval boundaries.
|
accepted seq_len range against interval boundaries.
|
||||||
@@ -1226,9 +1247,10 @@ class SchedulerBatchResultProcessor:
|
|||||||
interval = get_exec().mamba.mamba_track_interval
|
interval = get_exec().mamba.mamba_track_interval
|
||||||
|
|
||||||
if batch.spec_algorithm.is_none():
|
if batch.spec_algorithm.is_none():
|
||||||
self._mamba_assert_committed_len_lookahead(req)
|
lookahead = req.decode_batch_idx - batch.mamba_decode_batch_idx_cpu[i]
|
||||||
if req.kv_committed_len % interval == 0:
|
committed_len = req.kv_committed_len - lookahead
|
||||||
return True, req.kv_committed_len
|
if committed_len % interval == 0:
|
||||||
|
return True, committed_len
|
||||||
elif result.num_correct_drafts_per_req_cpu is not None:
|
elif result.num_correct_drafts_per_req_cpu is not None:
|
||||||
cur = req.seqlen - 1
|
cur = req.seqlen - 1
|
||||||
prev = cur - result.num_correct_drafts_per_req_cpu[i] - 1
|
prev = cur - result.num_correct_drafts_per_req_cpu[i] - 1
|
||||||
@@ -1237,19 +1259,13 @@ class SchedulerBatchResultProcessor:
|
|||||||
|
|
||||||
return False, 0
|
return False, 0
|
||||||
|
|
||||||
def mamba_lazy_post_decode_at_boundary(self, req: Req, batch: ScheduleBatch):
|
def mamba_lazy_post_decode_at_boundary(
|
||||||
"""Post-decode cleanup at a lazy-mode track boundary.
|
self, req: Req, batch: ScheduleBatch, track_idx: int
|
||||||
|
):
|
||||||
Finished reqs: if prealloc failed (other slot is -1), the forward
|
"""Commit a completed lazy-mode boundary and free its old slot."""
|
||||||
overwrote the only slot with corrupted state, so mark
|
req.mamba_last_track_idx = track_idx
|
||||||
is_insert=False to skip the cache insert. If the other slot is
|
req.mamba_next_track_idx = track_idx
|
||||||
occupied (stale prealloc from an overlap extra forward), free it
|
other_idx = 1 - track_idx
|
||||||
so the prealloc assert in the next prepare_for_decode holds.
|
|
||||||
|
|
||||||
Running reqs: free the old ping-pong slot so we go back to
|
|
||||||
holding only 1 slot until the next boundary.
|
|
||||||
"""
|
|
||||||
other_idx = 1 - req.mamba_next_track_idx
|
|
||||||
other_val = req.mamba_ping_pong_track_buffer[other_idx].item()
|
other_val = req.mamba_ping_pong_track_buffer[other_idx].item()
|
||||||
if other_val != -1:
|
if other_val != -1:
|
||||||
pool = batch.req_to_token_pool
|
pool = batch.req_to_token_pool
|
||||||
@@ -1257,5 +1273,3 @@ class SchedulerBatchResultProcessor:
|
|||||||
req.mamba_ping_pong_track_buffer[other_idx].unsqueeze(0)
|
req.mamba_ping_pong_track_buffer[other_idx].unsqueeze(0)
|
||||||
)
|
)
|
||||||
pool.set_mamba_ping_pong_slot(req, other_idx, -1)
|
pool.set_mamba_ping_pong_slot(req, other_idx, -1)
|
||||||
elif req.finished():
|
|
||||||
req.mamba_lazy_is_insert = False
|
|
||||||
|
|||||||
@@ -1388,14 +1388,8 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
return mamba_next_track_idx
|
return mamba_next_track_idx
|
||||||
|
|
||||||
def get_mamba_ping_pong_keep_idx(self, req: Req) -> int:
|
def get_mamba_ping_pong_keep_idx(self, req: Req) -> int:
|
||||||
"""Return the ping-pong index holding the most recent tracked state.
|
"""Return the ping-pong index holding the most recent tracked state."""
|
||||||
|
return req.mamba_last_track_idx
|
||||||
In lazy mode the valid state stays at next_track_idx (no eager swap).
|
|
||||||
In normal mode it is at the "other" index (swapped after each track).
|
|
||||||
"""
|
|
||||||
if self.enable_mamba_extra_buffer_lazy:
|
|
||||||
return req.mamba_next_track_idx
|
|
||||||
return self.get_mamba_ping_pong_other_idx(req.mamba_next_track_idx)
|
|
||||||
|
|
||||||
def _alloc_ping_pong_buffer(self, req: Req):
|
def _alloc_ping_pong_buffer(self, req: Req):
|
||||||
"""Allocate the ping-pong track buffer for a new request.
|
"""Allocate the ping-pong track buffer for a new request.
|
||||||
@@ -1422,6 +1416,11 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
buf[:n] = slots
|
buf[:n] = slots
|
||||||
req.mamba_ping_pong_track_buffer = buf
|
req.mamba_ping_pong_track_buffer = buf
|
||||||
req.mamba_next_track_idx = 0
|
req.mamba_next_track_idx = 0
|
||||||
|
req.mamba_last_track_idx = (
|
||||||
|
0
|
||||||
|
if self.enable_mamba_extra_buffer_lazy
|
||||||
|
else self.get_mamba_ping_pong_other_idx(0)
|
||||||
|
)
|
||||||
|
|
||||||
def set_mamba_ping_pong_slot(self, req: Req, idx: int, value):
|
def set_mamba_ping_pong_slot(self, req: Req, idx: int, value):
|
||||||
"""Update a ping-pong slot value and sync the device-side mapping.
|
"""Update a ping-pong slot value and sync the device-side mapping.
|
||||||
@@ -1442,8 +1441,6 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
|
|
||||||
Returns the old slot index (shape [1]) for cache insertion and
|
Returns the old slot index (shape [1]) for cache insertion and
|
||||||
replaces it with new_slot so the request can continue tracking.
|
replaces it with new_slot so the request can continue tracking.
|
||||||
In lazy mode the valid state is at next_track_idx; in normal mode
|
|
||||||
it is at the "other" index.
|
|
||||||
"""
|
"""
|
||||||
donate_idx = self.get_mamba_ping_pong_keep_idx(req)
|
donate_idx = self.get_mamba_ping_pong_keep_idx(req)
|
||||||
mamba_value_donated = (
|
mamba_value_donated = (
|
||||||
@@ -1512,6 +1509,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
# tensor on the req side while the new pool slot leaks).
|
# tensor on the req side while the new pool slot leaks).
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.mamba_next_track_idx = None
|
||||||
|
req.mamba_last_track_idx = None
|
||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
logger.info("Reset HybridReqToTokenPool")
|
logger.info("Reset HybridReqToTokenPool")
|
||||||
|
|||||||
@@ -60,6 +60,7 @@ class SessionSlot:
|
|||||||
mamba_pool_idx: Any = None
|
mamba_pool_idx: Any = None
|
||||||
mamba_ping_pong_track_buffer: Any = None
|
mamba_ping_pong_track_buffer: Any = None
|
||||||
mamba_next_track_idx: Any = None
|
mamba_next_track_idx: Any = None
|
||||||
|
mamba_last_track_idx: Any = None
|
||||||
mamba_last_track_seqlen: Any = None
|
mamba_last_track_seqlen: Any = None
|
||||||
mamba_branching_seqlen: Any = None
|
mamba_branching_seqlen: Any = None
|
||||||
|
|
||||||
@@ -83,6 +84,7 @@ class SessionSlot:
|
|||||||
self.mamba_pool_idx = req.mamba_pool_idx
|
self.mamba_pool_idx = req.mamba_pool_idx
|
||||||
self.mamba_ping_pong_track_buffer = req.mamba_ping_pong_track_buffer
|
self.mamba_ping_pong_track_buffer = req.mamba_ping_pong_track_buffer
|
||||||
self.mamba_next_track_idx = req.mamba_next_track_idx
|
self.mamba_next_track_idx = req.mamba_next_track_idx
|
||||||
|
self.mamba_last_track_idx = req.mamba_last_track_idx
|
||||||
self.mamba_last_track_seqlen = req.mamba_last_track_seqlen
|
self.mamba_last_track_seqlen = req.mamba_last_track_seqlen
|
||||||
self.mamba_branching_seqlen = req.mamba_branching_seqlen
|
self.mamba_branching_seqlen = req.mamba_branching_seqlen
|
||||||
|
|
||||||
@@ -99,6 +101,7 @@ class SessionSlot:
|
|||||||
req.mamba_pool_idx = None
|
req.mamba_pool_idx = None
|
||||||
req.mamba_ping_pong_track_buffer = None
|
req.mamba_ping_pong_track_buffer = None
|
||||||
req.mamba_next_track_idx = None
|
req.mamba_next_track_idx = None
|
||||||
|
req.mamba_last_track_idx = None
|
||||||
req.mamba_last_track_seqlen = None
|
req.mamba_last_track_seqlen = None
|
||||||
req.mamba_branching_seqlen = None
|
req.mamba_branching_seqlen = None
|
||||||
|
|
||||||
@@ -113,6 +116,7 @@ class SessionSlot:
|
|||||||
req.mamba_pool_idx = self.mamba_pool_idx
|
req.mamba_pool_idx = self.mamba_pool_idx
|
||||||
req.mamba_ping_pong_track_buffer = self.mamba_ping_pong_track_buffer
|
req.mamba_ping_pong_track_buffer = self.mamba_ping_pong_track_buffer
|
||||||
req.mamba_next_track_idx = self.mamba_next_track_idx
|
req.mamba_next_track_idx = self.mamba_next_track_idx
|
||||||
|
req.mamba_last_track_idx = self.mamba_last_track_idx
|
||||||
req.mamba_last_track_seqlen = self.mamba_last_track_seqlen
|
req.mamba_last_track_seqlen = self.mamba_last_track_seqlen
|
||||||
req.mamba_branching_seqlen = self.mamba_branching_seqlen
|
req.mamba_branching_seqlen = self.mamba_branching_seqlen
|
||||||
|
|
||||||
|
|||||||
@@ -110,9 +110,9 @@ class TestUnifiedMambaHiCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
|||||||
"--hicache-write-policy",
|
"--hicache-write-policy",
|
||||||
"write_through",
|
"write_through",
|
||||||
"--hicache-io-backend",
|
"--hicache-io-backend",
|
||||||
"direct",
|
"kernel",
|
||||||
"--hicache-mem-layout",
|
"--hicache-mem-layout",
|
||||||
"page_first_direct",
|
"page_first",
|
||||||
"--max-total-tokens",
|
"--max-total-tokens",
|
||||||
"12000",
|
"12000",
|
||||||
"--max-mamba-cache-size",
|
"--max-mamba-cache-size",
|
||||||
@@ -171,9 +171,9 @@ class TestUnifiedMambaHiCacheL3(AccuracyTwoPassMixin, CustomTestCase):
|
|||||||
"--hicache-storage-prefetch-policy",
|
"--hicache-storage-prefetch-policy",
|
||||||
"wait_complete",
|
"wait_complete",
|
||||||
"--hicache-io-backend",
|
"--hicache-io-backend",
|
||||||
"direct",
|
"kernel",
|
||||||
"--hicache-mem-layout",
|
"--hicache-mem-layout",
|
||||||
"page_first_direct",
|
"page_first",
|
||||||
"--hicache-storage-backend",
|
"--hicache-storage-backend",
|
||||||
"file",
|
"file",
|
||||||
"--max-mamba-cache-size",
|
"--max-mamba-cache-size",
|
||||||
|
|||||||
Reference in New Issue
Block a user