[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.allocator import TokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
|
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(
|
init_distributed_environment(
|
||||||
world_size=1,
|
world_size=1,
|
||||||
|
|||||||
@@ -18,10 +18,8 @@ from sglang.srt.mem_cache.memory_pool import (
|
|||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
ReqToTokenPool,
|
ReqToTokenPool,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||||
MLATokenToKVPoolHost,
|
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
|
||||||
get_mha_host_pool_cls,
|
|
||||||
)
|
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils.common import ceil_align
|
from sglang.srt.utils.common import ceil_align
|
||||||
|
|
||||||
|
|||||||
@@ -47,10 +47,8 @@ from sglang.srt.mem_cache.memory_pool import (
|
|||||||
MiniMaxSparseKVPool,
|
MiniMaxSparseKVPool,
|
||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||||
MLATokenToKVPoolHost,
|
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
|
||||||
get_mha_host_pool_cls,
|
|
||||||
)
|
|
||||||
from sglang.srt.mem_cache.radix_cache import (
|
from sglang.srt.mem_cache.radix_cache import (
|
||||||
RadixCache,
|
RadixCache,
|
||||||
RadixKey,
|
RadixKey,
|
||||||
|
|||||||
@@ -19,9 +19,11 @@ from sglang.srt.mem_cache.memory_pool_host import (
|
|||||||
HostPoolGroup,
|
HostPoolGroup,
|
||||||
LogicalHostPool,
|
LogicalHostPool,
|
||||||
MambaPoolHost,
|
MambaPoolHost,
|
||||||
MHATokenToKOnlyPoolHost,
|
|
||||||
MLATokenToKVPoolHost,
|
MLATokenToKVPoolHost,
|
||||||
PoolEntry,
|
PoolEntry,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.pool_host.mha import (
|
||||||
|
MHATokenToKOnlyPoolHost,
|
||||||
get_mha_host_pool_cls,
|
get_mha_host_pool_cls,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.unified_cache_components import ComponentType
|
from sglang.srt.mem_cache.unified_cache_components import ComponentType
|
||||||
|
|||||||
@@ -89,10 +89,8 @@ def maybe_register_hicache_draft(
|
|||||||
MHATokenToKVPool,
|
MHATokenToKVPool,
|
||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||||
MLATokenToKVPoolHost,
|
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
|
||||||
get_mha_host_pool_cls,
|
|
||||||
)
|
|
||||||
|
|
||||||
pool = draft_kv_pool
|
pool = draft_kv_pool
|
||||||
if isinstance(pool, HybridLinearKVPool):
|
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).
|
# Host RAM to leave free when sizing HiCache pools (OS, other processes).
|
||||||
HICACHE_HOST_MEMORY_RESERVE_BYTES: int = 10 * (1024**3)
|
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:
|
def sync_fixed_hicache_size(size: int, host_size: int) -> int:
|
||||||
"""Sync fixed-size HiCache token capacity across PP ranks.
|
"""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.hicache_storage import HiCacheStorageConfig
|
||||||
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
|
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(
|
logging.basicConfig(
|
||||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
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.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 import MHATokenToKVPool, MLATokenToKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import MLATokenToKVPoolHost
|
||||||
MHATokenToKVPoolHost,
|
|
||||||
MLATokenToKVPoolHost,
|
|
||||||
)
|
|
||||||
from sglang.srt.mem_cache.pool_host.common import (
|
from sglang.srt.mem_cache.pool_host.common import (
|
||||||
ALLOC_MEMORY_FUNCS,
|
ALLOC_MEMORY_FUNCS,
|
||||||
alloc_with_pin_memory,
|
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.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from types import SimpleNamespace
|
|||||||
import pytest
|
import pytest
|
||||||
import torch
|
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
|
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")
|
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
|
import torch
|
||||||
|
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.pool_host.mha import (
|
||||||
AsymmetricMHATokenToKVPoolHost,
|
AsymmetricMHATokenToKVPoolHost,
|
||||||
MHATokenToKVPoolHost,
|
MHATokenToKVPoolHost,
|
||||||
get_mha_host_pool_cls,
|
get_mha_host_pool_cls,
|
||||||
@@ -88,7 +88,7 @@ class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
|||||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||||
|
|
||||||
with mock.patch(
|
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,
|
create=True,
|
||||||
) as transfer:
|
) as transfer:
|
||||||
host.load_to_device_per_layer(
|
host.load_to_device_per_layer(
|
||||||
@@ -119,7 +119,7 @@ class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
|||||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||||
|
|
||||||
with mock.patch(
|
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,
|
create=True,
|
||||||
) as transfer:
|
) as transfer:
|
||||||
host.backup_from_device_all_layer(
|
host.backup_from_device_all_layer(
|
||||||
@@ -146,7 +146,7 @@ class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
|||||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||||
|
|
||||||
with mock.patch(
|
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,
|
create=True,
|
||||||
) as transfer:
|
) as transfer:
|
||||||
host.load_to_device_per_layer(
|
host.load_to_device_per_layer(
|
||||||
@@ -175,7 +175,7 @@ class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase):
|
|||||||
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||||
|
|
||||||
with mock.patch(
|
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,
|
create=True,
|
||||||
) as transfer:
|
) as transfer:
|
||||||
host.backup_from_device_all_layer(
|
host.backup_from_device_all_layer(
|
||||||
|
|||||||
@@ -25,15 +25,16 @@ from sglang.srt.mem_cache.memory_pool_host import (
|
|||||||
HostPoolGroup,
|
HostPoolGroup,
|
||||||
LogicalHostPool,
|
LogicalHostPool,
|
||||||
MambaPoolHost,
|
MambaPoolHost,
|
||||||
MHATokenToKVPoolHost,
|
|
||||||
MLATokenToKVPoolHost,
|
MLATokenToKVPoolHost,
|
||||||
PoolEntry,
|
PoolEntry,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=3, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=3, suite="base-a-test-cpu")
|
||||||
|
|
||||||
MEMORY_POOL_HOST_MODULE = "sglang.srt.mem_cache.memory_pool_host"
|
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:
|
def _indices(start: int, end: int) -> torch.Tensor:
|
||||||
@@ -230,21 +231,21 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
|
|||||||
|
|
||||||
with (
|
with (
|
||||||
mock.patch(
|
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(
|
side_effect=lambda **kwargs: _cpu_staged_mha_lf_pf_copy(
|
||||||
src_registry, **kwargs
|
src_registry, **kwargs
|
||||||
),
|
),
|
||||||
) as staged,
|
) as staged,
|
||||||
mock.patch(
|
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,
|
create=True,
|
||||||
) as fallback,
|
) as fallback,
|
||||||
mock.patch(
|
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,
|
side_effect=_cpu_jit_one_layer_mha_copy,
|
||||||
) as load,
|
) as load,
|
||||||
mock.patch(
|
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,
|
return_value=True,
|
||||||
) as can_use_write_back_jit_kernel,
|
) as can_use_write_back_jit_kernel,
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import unittest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
|
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.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
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 import MiniMaxSparseKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
HICACHE_HOST_MEMORY_RESERVE_BYTES,
|
HICACHE_HOST_MEMORY_RESERVE_BYTES,
|
||||||
MHATokenToKOnlyPoolHost,
|
|
||||||
MHATokenToKVPoolHost,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.pool_host.common import (
|
from sglang.srt.mem_cache.pool_host.common import (
|
||||||
ALLOC_MEMORY_FUNCS,
|
ALLOC_MEMORY_FUNCS,
|
||||||
alloc_with_pin_memory,
|
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.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user