diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 9448d77fc..6939a08d7 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -273,11 +273,8 @@ class MambaPool: def clear(self): # Zero the entire mamba cache before resetting free_slots # This ensures that when slots are reallocated, they start with clean state - if self.is_kda_cache: - for i in range(len(self.mamba_cache.conv)): - self.mamba_cache.conv[i].zero_() - else: - self.mamba_cache.conv.zero_() + 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)