diff --git a/python/sglang/srt/mem_cache/deepseek_v4_compress_state.py b/python/sglang/srt/mem_cache/deepseek_v4_compress_state.py index dff122578..bb0f703e3 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_compress_state.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_compress_state.py @@ -153,30 +153,25 @@ class CompressStatePool: allocation boilerplate. Requires ``self._size`` and ``self.last_dim`` to be set already. """ - if _is_hip: - self.kv_score_buffer = KVAndScore( - torch.empty((self._size, self.last_dim), dtype=dtype, device=device) - ) - else: - self.memory_saver_adapter = TorchMemorySaverAdapter.create( - enable=enable_memory_saver - ) - self.enable_custom_mem_pool, self.custom_mem_pool, _ = ( - maybe_init_custom_mem_pool(device=device) - ) - with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): - with ( - torch.cuda.use_mem_pool(self.custom_mem_pool) - if self.custom_mem_pool - else nullcontext() - ): - self.kv_score_buffer = KVAndScore( - torch.empty( - (self._size, self.last_dim), - dtype=dtype, - device=device, - ) + self.memory_saver_adapter = TorchMemorySaverAdapter.create( + enable=enable_memory_saver + ) + self.enable_custom_mem_pool, self.custom_mem_pool, _ = ( + maybe_init_custom_mem_pool(device=device) + ) + with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): + with ( + torch.cuda.use_mem_pool(self.custom_mem_pool) + if self.custom_mem_pool + else nullcontext() + ): + self.kv_score_buffer = KVAndScore( + torch.empty( + (self._size, self.last_dim), + dtype=dtype, + device=device, ) + ) @property def state_cache_3d(self) -> torch.Tensor: diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index 14c9e52cb..d1efe6db8 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -646,10 +646,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self._init_compressed_layer_mapping() - if _is_hip: - self._init_paged_compress_states(False) - else: - self._init_paged_compress_states(enable_memory_saver) + self._init_paged_compress_states(enable_memory_saver) def get_unified_kv(self, layer_id: int) -> torch.Tensor: # Under HiCache the compressed region is loaded H->D per layer; wait for this