From 70df09b83363e0127b43c83a6007d3938f815b2d Mon Sep 17 00:00:00 2001 From: iridiumine <42236072+iridiumine@users.noreply.github.com> Date: Thu, 2 Jul 2026 09:26:55 +0800 Subject: [PATCH] Fix MambaPool.clear_slots OOM by replacing expand-based tensor allocation with scalar zeroing (#27923) --- python/sglang/srt/mem_cache/memory_pool.py | 23 ++++++++++++++-------- 1 file changed, 15 insertions(+), 8 deletions(-) diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index e7d522728..82c9c98d2 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -667,18 +667,25 @@ class MambaPool: def clear_slots(self, indices: torch.Tensor): """Zero out mamba state at the given pool indices. Must run on forward stream.""" - need_size = len(indices) - for i in range(len(self.mamba_cache.conv)): - t = self.mamba_cache.conv[i] + if not _is_npu: + need_size = len(indices) + for i in range(len(self.mamba_cache.conv)): + t = self.mamba_cache.conv[i] + z = torch.zeros(1, dtype=t.dtype, device=t.device).expand( + t.shape[0], need_size, *t.shape[2:] + ) + t[:, indices] = z + t = self.mamba_cache.temporal z = torch.zeros(1, dtype=t.dtype, device=t.device).expand( t.shape[0], need_size, *t.shape[2:] ) t[:, indices] = z - t = self.mamba_cache.temporal - z = torch.zeros(1, dtype=t.dtype, device=t.device).expand( - t.shape[0], need_size, *t.shape[2:] - ) - t[:, indices] = z + else: + for i in range(len(self.mamba_cache.conv)): + t = self.mamba_cache.conv[i] + t[:, indices] = 0 + t = self.mamba_cache.temporal + t[:, indices] = 0 def copy_from(self, src_indices: torch.Tensor, dst_indices: torch.Tensor): """Clone mamba state (conv + temporal) from src slots into dst slots.