[Mamba] Support configurable conv-window layouts (#31059)
This commit is contained in:
@@ -311,6 +311,10 @@ class ReqToTokenPool:
|
|||||||
|
|
||||||
|
|
||||||
class MambaPool:
|
class MambaPool:
|
||||||
|
# Axis of each two-dimensional conv state that represents the sliding window.
|
||||||
|
# Upstream states use (dim, K-1); subclasses may preserve another layout.
|
||||||
|
conv_window_axis = -1
|
||||||
|
|
||||||
@dataclass(frozen=True, kw_only=True)
|
@dataclass(frozen=True, kw_only=True)
|
||||||
class State:
|
class State:
|
||||||
conv: List[torch.Tensor]
|
conv: List[torch.Tensor]
|
||||||
@@ -352,6 +356,46 @@ class MambaPool:
|
|||||||
intermediate_ssm: torch.Tensor
|
intermediate_ssm: torch.Tensor
|
||||||
intermediate_conv_window: List[torch.Tensor]
|
intermediate_conv_window: List[torch.Tensor]
|
||||||
|
|
||||||
|
def _allocate_deduplicated_conv_window(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
conv_shape: Tuple[int, int],
|
||||||
|
num_mamba_layers: int,
|
||||||
|
spec_state_size: int,
|
||||||
|
speculative_num_draft_tokens: int,
|
||||||
|
conv_dtype: torch.dtype,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
window_axis = self.conv_window_axis % len(conv_shape)
|
||||||
|
win = conv_shape[window_axis]
|
||||||
|
physical_conv_shape = list(conv_shape)
|
||||||
|
physical_conv_shape[window_axis] = speculative_num_draft_tokens + win - 1
|
||||||
|
phys = torch.zeros(
|
||||||
|
(
|
||||||
|
num_mamba_layers,
|
||||||
|
spec_state_size + 1,
|
||||||
|
*physical_conv_shape,
|
||||||
|
),
|
||||||
|
dtype=conv_dtype,
|
||||||
|
device="cuda",
|
||||||
|
)
|
||||||
|
physical_conv_strides = phys.stride()[2:]
|
||||||
|
window_stride = physical_conv_strides[window_axis]
|
||||||
|
view = phys.as_strided(
|
||||||
|
(
|
||||||
|
phys.shape[0],
|
||||||
|
phys.shape[1],
|
||||||
|
speculative_num_draft_tokens,
|
||||||
|
*conv_shape,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
phys.stride(0),
|
||||||
|
phys.stride(1),
|
||||||
|
window_stride,
|
||||||
|
*physical_conv_strides,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return phys, view
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -538,36 +582,12 @@ class MambaPool:
|
|||||||
if dedup_conv_window:
|
if dedup_conv_window:
|
||||||
intermediate_conv_window_cache = []
|
intermediate_conv_window_cache = []
|
||||||
for conv_shape in conv_state_shape:
|
for conv_shape in conv_state_shape:
|
||||||
conv_dim, win = conv_shape # win == conv_kernel - 1 == K-1
|
phys, view = self._allocate_deduplicated_conv_window(
|
||||||
shared_win = (
|
conv_shape=conv_shape,
|
||||||
speculative_num_draft_tokens + win - 1
|
num_mamba_layers=num_mamba_layers,
|
||||||
) # D + (K-1) - 1
|
spec_state_size=spec_state_size,
|
||||||
phys = torch.zeros(
|
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||||||
size=(
|
conv_dtype=conv_dtype,
|
||||||
num_mamba_layers,
|
|
||||||
spec_state_size + 1,
|
|
||||||
conv_dim,
|
|
||||||
shared_win,
|
|
||||||
),
|
|
||||||
dtype=conv_dtype,
|
|
||||||
device="cuda",
|
|
||||||
)
|
|
||||||
# view[l, s, step, d, w] = phys[l, s, d, step + w]
|
|
||||||
view = phys.as_strided(
|
|
||||||
(
|
|
||||||
phys.shape[0],
|
|
||||||
phys.shape[1],
|
|
||||||
speculative_num_draft_tokens,
|
|
||||||
conv_dim,
|
|
||||||
win,
|
|
||||||
),
|
|
||||||
(
|
|
||||||
phys.stride(0),
|
|
||||||
phys.stride(1),
|
|
||||||
phys.stride(3), # step -> shared-win axis (stride 1)
|
|
||||||
phys.stride(2), # dim
|
|
||||||
phys.stride(3), # win -> shared-win axis (stride 1)
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
self._intermediate_conv_window_phys.append(phys)
|
self._intermediate_conv_window_phys.append(phys)
|
||||||
intermediate_conv_window_cache.append(view)
|
intermediate_conv_window_cache.append(view)
|
||||||
@@ -818,6 +838,8 @@ class MambaPool:
|
|||||||
class HybridReqToTokenPool(ReqToTokenPool):
|
class HybridReqToTokenPool(ReqToTokenPool):
|
||||||
"""A memory pool that maps a request to its token locations."""
|
"""A memory pool that maps a request to its token locations."""
|
||||||
|
|
||||||
|
mamba_pool_cls = MambaPool
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -880,7 +902,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
linear_replayssm_cache_len: int = 16,
|
linear_replayssm_cache_len: int = 16,
|
||||||
mamba_envelope_layout: bool = False,
|
mamba_envelope_layout: bool = False,
|
||||||
):
|
):
|
||||||
self.mamba_pool = MambaPool(
|
self.mamba_pool = self.mamba_pool_cls(
|
||||||
size=mamba_size,
|
size=mamba_size,
|
||||||
spec_state_size=mamba_spec_state_size,
|
spec_state_size=mamba_spec_state_size,
|
||||||
cache_params=cache_params,
|
cache_params=cache_params,
|
||||||
|
|||||||
@@ -18,7 +18,11 @@ from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
|||||||
from sglang.srt.mem_cache.common import available_and_evictable_str
|
from sglang.srt.mem_cache.common import available_and_evictable_str
|
||||||
from sglang.srt.mem_cache.hi_mamba_radix_cache import HiMambaRadixCache
|
from sglang.srt.mem_cache.hi_mamba_radix_cache import HiMambaRadixCache
|
||||||
from sglang.srt.mem_cache.mamba_radix_cache import LRUList, MambaRadixCache, TreeNode
|
from sglang.srt.mem_cache.mamba_radix_cache import LRUList, MambaRadixCache, TreeNode
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import (
|
||||||
|
HybridLinearKVPool,
|
||||||
|
HybridReqToTokenPool,
|
||||||
|
MambaPool,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixKey
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||||
@@ -154,6 +158,61 @@ class TestMamba(unittest.TestCase):
|
|||||||
req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1
|
req_to_token_pool.mamba_allocator.available_size() == mamba_cache_size - 1
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_mamba_pool_deduplicated_conv_window_axis(self):
|
||||||
|
class WindowFirstMambaPool(MambaPool):
|
||||||
|
conv_window_axis = 0
|
||||||
|
|
||||||
|
num_mamba_layers = 2
|
||||||
|
spec_state_size = 3
|
||||||
|
speculative_num_draft_tokens = 4
|
||||||
|
window_size = 3
|
||||||
|
conv_dim = 5
|
||||||
|
|
||||||
|
pool = object.__new__(WindowFirstMambaPool)
|
||||||
|
physical, view = pool._allocate_deduplicated_conv_window(
|
||||||
|
conv_shape=(window_size, conv_dim),
|
||||||
|
num_mamba_layers=num_mamba_layers,
|
||||||
|
spec_state_size=spec_state_size,
|
||||||
|
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||||||
|
conv_dtype=torch.float32,
|
||||||
|
)
|
||||||
|
|
||||||
|
shared_window_size = speculative_num_draft_tokens + window_size - 1
|
||||||
|
self.assertEqual(
|
||||||
|
physical.shape,
|
||||||
|
(
|
||||||
|
num_mamba_layers,
|
||||||
|
spec_state_size + 1,
|
||||||
|
shared_window_size,
|
||||||
|
conv_dim,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
view.shape,
|
||||||
|
(
|
||||||
|
num_mamba_layers,
|
||||||
|
spec_state_size + 1,
|
||||||
|
speculative_num_draft_tokens,
|
||||||
|
window_size,
|
||||||
|
conv_dim,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
physical.copy_(
|
||||||
|
torch.arange(
|
||||||
|
physical.numel(), dtype=physical.dtype, device=physical.device
|
||||||
|
).reshape_as(physical)
|
||||||
|
)
|
||||||
|
for step in range(speculative_num_draft_tokens):
|
||||||
|
torch.testing.assert_close(
|
||||||
|
view[:, :, step],
|
||||||
|
physical[:, :, step : step + window_size],
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(view[:, :, :-1, 1:], view[:, :, 1:, :-1])
|
||||||
|
|
||||||
|
view[0, 0, 0, 1, 0] = -1
|
||||||
|
self.assertEqual(view[0, 0, 1, 0, 0].item(), -1)
|
||||||
|
|
||||||
def test_mamba_radix_cache_1(self):
|
def test_mamba_radix_cache_1(self):
|
||||||
tree, allocator, req_to_token_pool, make_dummy_req = (
|
tree, allocator, req_to_token_pool, make_dummy_req = (
|
||||||
self._setup_tree_and_allocator()
|
self._setup_tree_and_allocator()
|
||||||
|
|||||||
Reference in New Issue
Block a user