[mem_cache][7/N] refactor: move MLATokenToKVPoolHost to pool_host.mla (#30616)
This commit is contained in:
@@ -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(),
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user