[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:
co-authored by
Zhiyao Jiang
parent
c32c4ef79c
commit
f5bcd00e16
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user