[mem_cache][6/N] refactor: move MHA host-pool into pool_host/mha.py (#30249)

This commit is contained in:
shuwenn
2026-07-08 20:12:52 +08:00
committed by GitHub
parent 4c5fe42be4
commit 108a183f6b
15 changed files with 1266 additions and 1222 deletions
+2 -4
View File
@@ -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