[AMD] DSv4: bring HIP compress-state pool into the memory_saver KV_CACHE region (#31747)

Co-authored-by: Zhiyao Jiang <jessicajiang324@gmail.com>
This commit is contained in:
Xinyu Jiang
2026-07-29 02:41:08 -07:00
committed by GitHub
co-authored by Zhiyao Jiang
parent c32c4ef79c
commit f5bcd00e16
2 changed files with 19 additions and 27 deletions
@@ -153,30 +153,25 @@ class CompressStatePool:
allocation boilerplate. Requires ``self._size`` and ``self.last_dim`` allocation boilerplate. Requires ``self._size`` and ``self.last_dim``
to be set already. to be set already.
""" """
if _is_hip: self.memory_saver_adapter = TorchMemorySaverAdapter.create(
self.kv_score_buffer = KVAndScore( enable=enable_memory_saver
torch.empty((self._size, self.last_dim), dtype=dtype, device=device) )
) self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
else: maybe_init_custom_mem_pool(device=device)
self.memory_saver_adapter = TorchMemorySaverAdapter.create( )
enable=enable_memory_saver with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
) with (
self.enable_custom_mem_pool, self.custom_mem_pool, _ = ( torch.cuda.use_mem_pool(self.custom_mem_pool)
maybe_init_custom_mem_pool(device=device) if self.custom_mem_pool
) else nullcontext()
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): ):
with ( self.kv_score_buffer = KVAndScore(
torch.cuda.use_mem_pool(self.custom_mem_pool) torch.empty(
if self.custom_mem_pool (self._size, self.last_dim),
else nullcontext() dtype=dtype,
): device=device,
self.kv_score_buffer = KVAndScore(
torch.empty(
(self._size, self.last_dim),
dtype=dtype,
device=device,
)
) )
)
@property @property
def state_cache_3d(self) -> torch.Tensor: def state_cache_3d(self) -> torch.Tensor:
@@ -646,10 +646,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self._init_compressed_layer_mapping() self._init_compressed_layer_mapping()
if _is_hip: self._init_paged_compress_states(enable_memory_saver)
self._init_paged_compress_states(False)
else:
self._init_paged_compress_states(enable_memory_saver)
def get_unified_kv(self, layer_id: int) -> torch.Tensor: def get_unified_kv(self, layer_id: int) -> torch.Tensor:
# Under HiCache the compressed region is loaded H->D per layer; wait for this # Under HiCache the compressed region is loaded H->D per layer; wait for this