diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index ef9df64e4..9d1f0450f 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2583,7 +2583,10 @@ class Scheduler( release_kv_cache(req, self.tree_cache) # For mamba radix cache - if req.mamba_pool_idx is not None: + if ( + req.mamba_pool_idx is not None + and self.disaggregation_mode != DisaggregationMode.DECODE + ): release_kv_cache(req, self.tree_cache, is_insert=False) logger.debug(f"Abort queued request. {req.rid=}") diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 11446e102..b75d3be86 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -276,6 +276,11 @@ class MambaPool: self.mamba_cache.conv[i][:, select_index] = 0 self.mamba_cache.temporal[:, select_index] = 0 + # fill allocated slots with zeros + for i in range(len(self.mamba_cache.conv)): + self.mamba_cache.conv[i][:, select_index] = 0 + self.mamba_cache.temporal[:, select_index] = 0 + return select_index def free(self, free_index: torch.Tensor): @@ -284,12 +289,6 @@ class MambaPool: self.free_slots = torch.cat((self.free_slots, free_index)) def clear(self): - # Zero the entire mamba cache before resetting free_slots - # This ensures that when slots are reallocated, they start with clean state - for i in range(len(self.mamba_cache.conv)): - self.mamba_cache.conv[i].zero_() - self.mamba_cache.temporal.zero_() - self.free_slots = torch.arange(self.size, dtype=torch.int64, device=self.device) def copy_from(self, src_index: torch.Tensor, dst_index: torch.Tensor):