[mem_cache][7/N] refactor: move MLATokenToKVPoolHost to pool_host.mla (#30616)

This commit is contained in:
shuwenn
2026-07-13 14:23:38 +08:00
committed by GitHub
parent cbcbef6811
commit 9dd57ef8c4
13 changed files with 602 additions and 564 deletions
+1 -1
View File
@@ -5,12 +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 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.mem_cache.pool_host.mla import MLATokenToKVPoolHost
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
from sglang.test.ci.ci_register import register_cuda_ci
@@ -15,12 +15,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 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.mem_cache.pool_host.mla import MLATokenToKVPoolHost
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
@@ -4,14 +4,12 @@ import unittest
import torch
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
from sglang.srt.mem_cache.memory_pool_host import (
DSAIndexerPoolHost,
MLATokenToKVPoolHost,
)
from sglang.srt.mem_cache.memory_pool_host import DSAIndexerPoolHost
from sglang.srt.mem_cache.pool_host.common import (
ALLOC_MEMORY_FUNCS,
alloc_with_pin_memory,
)
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
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
@@ -148,7 +146,7 @@ class TestDSAHiCacheTransfer(unittest.TestCase):
@unittest.skipIf(
is_hip(),
'`io_backend="kernel"` path in memory_pool_host.backup_from_device_all_layer '
'`io_backend="kernel"` path in MLATokenToKVPoolHost.backup_from_device_all_layer '
"raises ValueError on AMD (only the `direct` IO backend is wired for ROCm). "
"The other 62 tests in this file pass on AMD.",
)
@@ -25,16 +25,17 @@ from sglang.srt.mem_cache.memory_pool_host import (
HostPoolGroup,
LogicalHostPool,
MambaPoolHost,
MLATokenToKVPoolHost,
PoolEntry,
)
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
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"
MLA_POOL_HOST_MODULE = "sglang.srt.mem_cache.pool_host.mla"
def _indices(start: int, end: int) -> torch.Tensor:
@@ -155,7 +156,7 @@ class _FakeDeviceModule:
class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
def _patched_transfers(self, src_registry=None):
def _patched_transfers(self, src_registry=None, module=MEMORY_POOL_HOST_MODULE):
staged_side_effect = None
if src_registry is not None:
staged_side_effect = lambda **kwargs: _cpu_staged_lf_pf_copy(
@@ -163,15 +164,15 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
)
return (
mock.patch(
f"{MEMORY_POOL_HOST_MODULE}.jit_transfer_hicache_all_layer_mla_staged_lf_pf",
f"{module}.jit_transfer_hicache_all_layer_mla_staged_lf_pf",
side_effect=staged_side_effect,
),
mock.patch(
f"{MEMORY_POOL_HOST_MODULE}.transfer_kv_all_layer_mla_lf_pf",
f"{module}.transfer_kv_all_layer_mla_lf_pf",
create=True,
),
mock.patch(
f"{MEMORY_POOL_HOST_MODULE}.transfer_kv_per_layer_mla_pf_lf",
f"{module}.transfer_kv_per_layer_mla_pf_lf",
side_effect=_cpu_per_layer_pf_lf_copy,
create=True,
),
@@ -329,16 +330,18 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
)
src_registry = {_ptr_key_from_layers(device_layers): device_layers}
staged_patch, fallback_patch, _ = self._patched_transfers(src_registry)
staged_patch, fallback_patch, _ = self._patched_transfers(
src_registry, module=MLA_POOL_HOST_MODULE
)
with (
staged_patch as staged,
fallback_patch as fallback,
mock.patch(
f"{MEMORY_POOL_HOST_MODULE}.jit_transfer_hicache_one_layer_mla",
f"{MLA_POOL_HOST_MODULE}.jit_transfer_hicache_one_layer_mla",
side_effect=_cpu_jit_one_layer_mla_copy,
) as load,
mock.patch(
f"{MEMORY_POOL_HOST_MODULE}.can_use_write_back_jit_kernel",
f"{MLA_POOL_HOST_MODULE}.can_use_write_back_jit_kernel",
return_value=True,
) as can_use_write_back_jit_kernel,
):
@@ -42,14 +42,14 @@ def _fake_mooncake_modules(fake_store_cls, replicate_config_cls):
}
def _fake_memory_pool_host_module():
memory_pool_host = types.ModuleType("sglang.srt.mem_cache.memory_pool_host")
def _fake_pool_host_mla_module():
pool_host_mla = types.ModuleType("sglang.srt.mem_cache.pool_host.mla")
class MLATokenToKVPoolHost:
pass
memory_pool_host.MLATokenToKVPoolHost = MLATokenToKVPoolHost
return memory_pool_host
pool_host_mla.MLATokenToKVPoolHost = MLATokenToKVPoolHost
return pool_host_mla
def _fake_pool_host_module():
@@ -68,8 +68,8 @@ def _fake_pool_host_module():
def _fake_host_pool_modules():
return {
"sglang.srt.mem_cache.memory_pool_host": _fake_memory_pool_host_module(),
"sglang.srt.mem_cache.pool_host": _fake_pool_host_module(),
"sglang.srt.mem_cache.pool_host.mla": _fake_pool_host_mla_module(),
}