[mem_cache][6/N] refactor: move MHA host-pool into pool_host/mha.py (#30249)
This commit is contained in:
@@ -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