[mem_cache][6/N] refactor: move MHA host-pool into pool_host/mha.py (#30249)
This commit is contained in:
@@ -16,7 +16,7 @@ from sglang.srt.managers.cache_controller import (
|
||||
)
|
||||
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
|
||||
from sglang.srt.mem_cache.memory_pool_host import MHATokenToKVPoolHost
|
||||
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||
|
||||
init_distributed_environment(
|
||||
world_size=1,
|
||||
|
||||
@@ -18,10 +18,8 @@ from sglang.srt.mem_cache.memory_pool import (
|
||||
MLATokenToKVPool,
|
||||
ReqToTokenPool,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
MLATokenToKVPoolHost,
|
||||
get_mha_host_pool_cls,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.common import ceil_align
|
||||
|
||||
|
||||
@@ -47,10 +47,8 @@ from sglang.srt.mem_cache.memory_pool import (
|
||||
MiniMaxSparseKVPool,
|
||||
MLATokenToKVPool,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
MLATokenToKVPoolHost,
|
||||
get_mha_host_pool_cls,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
|
||||
from sglang.srt.mem_cache.radix_cache import (
|
||||
RadixCache,
|
||||
RadixKey,
|
||||
|
||||
@@ -19,9 +19,11 @@ from sglang.srt.mem_cache.memory_pool_host import (
|
||||
HostPoolGroup,
|
||||
LogicalHostPool,
|
||||
MambaPoolHost,
|
||||
MHATokenToKOnlyPoolHost,
|
||||
MLATokenToKVPoolHost,
|
||||
PoolEntry,
|
||||
)
|
||||
from sglang.srt.mem_cache.pool_host.mha import (
|
||||
MHATokenToKOnlyPoolHost,
|
||||
get_mha_host_pool_cls,
|
||||
)
|
||||
from sglang.srt.mem_cache.unified_cache_components import ComponentType
|
||||
|
||||
@@ -89,10 +89,8 @@ def maybe_register_hicache_draft(
|
||||
MHATokenToKVPool,
|
||||
MLATokenToKVPool,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
MLATokenToKVPoolHost,
|
||||
get_mha_host_pool_cls,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
|
||||
|
||||
pool = draft_kv_pool
|
||||
if isinstance(pool, HybridLinearKVPool):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -24,6 +24,8 @@ _is_hip = is_hip()
|
||||
# Host RAM to leave free when sizing HiCache pools (OS, other processes).
|
||||
HICACHE_HOST_MEMORY_RESERVE_BYTES: int = 10 * (1024**3)
|
||||
|
||||
_WRITE_BACK_STAGING_PAGE_CHUNK = 64
|
||||
|
||||
|
||||
def sync_fixed_hicache_size(size: int, host_size: int) -> int:
|
||||
"""Sync fixed-size HiCache token capacity across PP ranks.
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -8,7 +8,7 @@ from aibrix_kvcache_storage import AibrixKVCacheStorage
|
||||
|
||||
from sglang.srt.mem_cache.hicache_storage import HiCacheStorageConfig
|
||||
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
|
||||
from sglang.srt.mem_cache.memory_pool_host import MHATokenToKVPoolHost
|
||||
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
||||
|
||||
@@ -5,14 +5,12 @@ import torch
|
||||
|
||||
from sglang.jit_kernel.hicache import can_use_write_back_jit_kernel
|
||||
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, MLATokenToKVPool
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
MHATokenToKVPoolHost,
|
||||
MLATokenToKVPoolHost,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||
from sglang.srt.mem_cache.pool_host.common import (
|
||||
ALLOC_MEMORY_FUNCS,
|
||||
alloc_with_pin_memory,
|
||||
)
|
||||
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ from types import SimpleNamespace
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.memory_pool_host import AsymmetricMHATokenToKVPoolHost
|
||||
from sglang.srt.mem_cache.pool_host.mha import AsymmetricMHATokenToKVPoolHost
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
@@ -6,7 +6,7 @@ from unittest import mock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
from sglang.srt.mem_cache.pool_host.mha import (
|
||||
AsymmetricMHATokenToKVPoolHost,
|
||||
MHATokenToKVPoolHost,
|
||||
get_mha_host_pool_cls,
|
||||
@@ -88,7 +88,7 @@ class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host.transfer_kv_per_layer_mla_pf_lf",
|
||||
"sglang.srt.mem_cache.pool_host.mha.transfer_kv_per_layer_mla_pf_lf",
|
||||
create=True,
|
||||
) as transfer:
|
||||
host.load_to_device_per_layer(
|
||||
@@ -119,7 +119,7 @@ class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host.transfer_kv_all_layer_mla_lf_pf",
|
||||
"sglang.srt.mem_cache.pool_host.mha.transfer_kv_all_layer_mla_lf_pf",
|
||||
create=True,
|
||||
) as transfer:
|
||||
host.backup_from_device_all_layer(
|
||||
@@ -146,7 +146,7 @@ class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host.transfer_kv_per_layer_direct_pf_lf",
|
||||
"sglang.srt.mem_cache.pool_host.mha.transfer_kv_per_layer_direct_pf_lf",
|
||||
create=True,
|
||||
) as transfer:
|
||||
host.load_to_device_per_layer(
|
||||
@@ -175,7 +175,7 @@ class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host.transfer_kv_all_layer_direct_lf_pf",
|
||||
"sglang.srt.mem_cache.pool_host.mha.transfer_kv_all_layer_direct_lf_pf",
|
||||
create=True,
|
||||
) as transfer:
|
||||
host.backup_from_device_all_layer(
|
||||
|
||||
@@ -25,15 +25,16 @@ from sglang.srt.mem_cache.memory_pool_host import (
|
||||
HostPoolGroup,
|
||||
LogicalHostPool,
|
||||
MambaPoolHost,
|
||||
MHATokenToKVPoolHost,
|
||||
MLATokenToKVPoolHost,
|
||||
PoolEntry,
|
||||
)
|
||||
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=3, suite="base-a-test-cpu")
|
||||
|
||||
MEMORY_POOL_HOST_MODULE = "sglang.srt.mem_cache.memory_pool_host"
|
||||
MHA_POOL_HOST_MODULE = "sglang.srt.mem_cache.pool_host.mha"
|
||||
|
||||
|
||||
def _indices(start: int, end: int) -> torch.Tensor:
|
||||
@@ -230,21 +231,21 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
|
||||
|
||||
with (
|
||||
mock.patch(
|
||||
f"{MEMORY_POOL_HOST_MODULE}.jit_transfer_hicache_all_layer_staged_lf_pf",
|
||||
f"{MHA_POOL_HOST_MODULE}.jit_transfer_hicache_all_layer_staged_lf_pf",
|
||||
side_effect=lambda **kwargs: _cpu_staged_mha_lf_pf_copy(
|
||||
src_registry, **kwargs
|
||||
),
|
||||
) as staged,
|
||||
mock.patch(
|
||||
f"{MEMORY_POOL_HOST_MODULE}.transfer_kv_all_layer_lf_pf",
|
||||
f"{MHA_POOL_HOST_MODULE}.transfer_kv_all_layer_lf_pf",
|
||||
create=True,
|
||||
) as fallback,
|
||||
mock.patch(
|
||||
f"{MEMORY_POOL_HOST_MODULE}.jit_transfer_hicache_one_layer",
|
||||
f"{MHA_POOL_HOST_MODULE}.jit_transfer_hicache_one_layer",
|
||||
side_effect=_cpu_jit_one_layer_mha_copy,
|
||||
) as load,
|
||||
mock.patch(
|
||||
f"{MEMORY_POOL_HOST_MODULE}.can_use_write_back_jit_kernel",
|
||||
f"{MHA_POOL_HOST_MODULE}.can_use_write_back_jit_kernel",
|
||||
return_value=True,
|
||||
) as can_use_write_back_jit_kernel,
|
||||
):
|
||||
|
||||
@@ -5,7 +5,7 @@ import unittest
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
|
||||
from sglang.srt.mem_cache.memory_pool_host import MHATokenToKVPoolHost
|
||||
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
|
||||
@@ -11,13 +11,15 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
||||
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
HICACHE_HOST_MEMORY_RESERVE_BYTES,
|
||||
MHATokenToKOnlyPoolHost,
|
||||
MHATokenToKVPoolHost,
|
||||
)
|
||||
from sglang.srt.mem_cache.pool_host.common import (
|
||||
ALLOC_MEMORY_FUNCS,
|
||||
alloc_with_pin_memory,
|
||||
)
|
||||
from sglang.srt.mem_cache.pool_host.mha import (
|
||||
MHATokenToKOnlyPoolHost,
|
||||
MHATokenToKVPoolHost,
|
||||
)
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||
|
||||
|
||||
Reference in New Issue
Block a user