Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 25 additions & 3 deletions python/sglang/srt/mem_cache/unified_memory_pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,11 @@ def _prod(iterable) -> int:


def _store_dtype_for(kv_cache_dtype: torch.dtype) -> torch.dtype:
if kv_cache_dtype in (torch.float8_e5m2, torch.float8_e4m3fn):
if kv_cache_dtype in (
torch.float8_e5m2,
torch.float8_e4m3fn,
torch.float8_e4m3fnuz,
):
return torch.uint8
return kv_cache_dtype

Expand Down Expand Up @@ -127,12 +131,21 @@ class MHASubPoolSpec(SubPoolSpec):
head_num: int
head_dim: int
store_dtype: torch.dtype
kv_cache_dtype: Optional[torch.dtype] = None
v_head_dim: Optional[int] = None

def __post_init__(self):
super().__post_init__()
assert self.head_num > 0, f"head_num must be positive; got {self.head_num}"
assert self.head_dim > 0, f"head_dim must be positive; got {self.head_dim}"
if self.kv_cache_dtype is None:
object.__setattr__(self, "kv_cache_dtype", self.store_dtype)
expected_store_dtype = _store_dtype_for(self.kv_cache_dtype)
assert self.store_dtype == expected_store_dtype, (
"MHASubPoolSpec.store_dtype must match the storage dtype required by "
f"kv_cache_dtype; got kv_cache_dtype={self.kv_cache_dtype}, "
f"store_dtype={self.store_dtype}, expected={expected_store_dtype}"
)
if self.v_head_dim is None:
object.__setattr__(self, "v_head_dim", self.head_dim)
assert self.v_head_dim > 0, (
Expand Down Expand Up @@ -569,7 +582,7 @@ def __init__(
super().__init__(
size=view_rows - page_size,
page_size=page_size,
dtype=spec.store_dtype,
dtype=spec.kv_cache_dtype,
head_num=spec.head_num,
head_dim=spec.head_dim,
layer_num=spec.layer_num,
Expand Down Expand Up @@ -1313,6 +1326,7 @@ def init_unified_mamba_pools(
head_num=head_num,
head_dim=head_dim,
store_dtype=store_dtype,
kv_cache_dtype=kv_cache_dtype,
grow_direction="down",
)
cp = mamba2_cache_params
Expand Down Expand Up @@ -1536,7 +1550,11 @@ def __init__(
"UnifiedSWAKVPool: full and swa sub-pools must share store_dtype; got "
f"full={full_spec.store_dtype}, swa={swa_spec.store_dtype}"
)
self.dtype = full_spec.store_dtype
assert full_spec.kv_cache_dtype == swa_spec.kv_cache_dtype, (
"UnifiedSWAKVPool: full and swa sub-pools must share kv_cache_dtype; got "
f"full={full_spec.kv_cache_dtype}, swa={swa_spec.kv_cache_dtype}"
)
self.dtype = full_spec.kv_cache_dtype
self.head_num = full_spec.head_num
self.head_dim = full_spec.head_dim
self.device = unified_buffer.device
Expand Down Expand Up @@ -1803,6 +1821,7 @@ def init_unified_swa_pools(
head_dim=head_dim,
v_head_dim=v_head_dim,
store_dtype=store_dtype,
kv_cache_dtype=kv_cache_dtype,
grow_direction="down",
)
swa_spec = MHASubPoolSpec(
Expand All @@ -1812,6 +1831,7 @@ def init_unified_swa_pools(
head_dim=swa_head_dim,
v_head_dim=swa_v_head_dim,
store_dtype=store_dtype,
kv_cache_dtype=kv_cache_dtype,
grow_direction="up",
)
legacy_allocator_capacities = {}
Expand Down Expand Up @@ -1994,6 +2014,7 @@ def init_unified_mamba_swa_pools(
head_dim=head_dim,
v_head_dim=v_head_dim,
store_dtype=store_dtype,
kv_cache_dtype=kv_cache_dtype,
grow_direction="down",
)
swa_spec = MHASubPoolSpec(
Expand All @@ -2003,6 +2024,7 @@ def init_unified_mamba_swa_pools(
head_dim=swa_head_dim,
v_head_dim=swa_v_head_dim,
store_dtype=store_dtype,
kv_cache_dtype=kv_cache_dtype,
grow_direction="float",
)
cp = mamba2_cache_params
Expand Down
31 changes: 31 additions & 0 deletions test/registered/unit/mem_cache/test_unified_mha_views.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@
_DTYPE = torch.bfloat16
_ITEM = _DTYPE.itemsize
_BLOCKS = 2 * _L
_HAS_FP8 = hasattr(torch, "float8_e4m3fn")


def _mha_spec(head_dim=_D, v_head_dim=None, layer_num=_L, grow="down"):
Expand Down Expand Up @@ -432,6 +433,36 @@ def test_factory_wires_matching_multipliers(self):
self.assertEqual(b.token_to_kv_pool.swa_kv_pool.k_buffer[0].dim(), 3)
self.assertGreater(pool.view_tail_pad_bytes, 0)

@unittest.skipUnless(_HAS_FP8, "requires torch.float8_e4m3fn")
def test_swa_factory_preserves_fp8_logical_dtype(self):
from sglang.srt.mem_cache.unified_memory_pool import init_unified_swa_pools

b = init_unified_swa_pools(
device="cpu",
kv_cache_dtype=torch.float8_e4m3fn,
head_num=2,
head_dim=8,
v_head_dim=8,
swa_head_num=2,
swa_head_dim=8,
swa_v_head_dim=8,
page_size=1,
start_layer=0,
end_layer=4,
swa_attention_layer_ids=[1, 3],
full_attention_layer_ids=[0, 2],
full_max_total_num_tokens=64,
swa_max_total_num_tokens=32,
enable_memory_saver=False,
need_sort=False,
)

self.assertEqual(b.token_to_kv_pool.dtype, torch.float8_e4m3fn)
for pool in (b.token_to_kv_pool.full_kv_pool, b.token_to_kv_pool.swa_kv_pool):
self.assertEqual(pool.dtype, torch.float8_e4m3fn)
self.assertEqual(pool.store_dtype, torch.uint8)
self.assertEqual(pool.k_buffer[0].dtype, torch.uint8)

def test_rebind_emits_kernel_facing_full_and_build_derives_swa(self):
"""rebind_write_loc rebinds out_cache_loc to FULL-kernel-facing ids, and
the SWA write loc is derived pointwise from those kernel-facing values."""
Expand Down
Loading