[unified-memory] PD disaggregation for every unified pool shape (#37506)
This commit is contained in:
@@ -82,8 +82,10 @@ class Memory(msgspec.Struct):
|
|||||||
"Replace the statically-partitioned hybrid-model pools (full-attn KV + "
|
"Replace the statically-partitioned hybrid-model pools (full-attn KV + "
|
||||||
"SWA/Mamba state) with one byte buffer split dynamically between "
|
"SWA/Mamba state) with one byte buffer split dynamically between "
|
||||||
"sub-pools. Requires the Triton attention / linear-attn / Mamba "
|
"sub-pools. Requires the Triton attention / linear-attn / Mamba "
|
||||||
"backends; not yet compatible with PD disaggregation or speculative "
|
"backends. PD disaggregation is supported over mooncake at equal "
|
||||||
"decoding.",
|
"attention TP with pp=1; not yet compatible with hierarchical / "
|
||||||
|
"host-tiered KV cache, prefill cuda-graph capture, or speculative "
|
||||||
|
"decoding other than DSPARK.",
|
||||||
] = False
|
] = False
|
||||||
enable_session_radix_cache: A[
|
enable_session_radix_cache: A[
|
||||||
bool,
|
bool,
|
||||||
|
|||||||
@@ -1416,7 +1416,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
decode_req.req.kv.req_pool_idx, window_start:seq_len
|
decode_req.req.kv.req_pool_idx, window_start:seq_len
|
||||||
]
|
]
|
||||||
window_kv_indices_swa = (
|
window_kv_indices_swa = (
|
||||||
self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
|
self.token_to_kv_pool_allocator.translate_swa_indices_for_transfer(
|
||||||
window_kv_indices_full
|
window_kv_indices_full
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -957,6 +957,13 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
prefill_data_indices=prefill_kv_indices,
|
prefill_data_indices=prefill_kv_indices,
|
||||||
dst_data_indices=dst_kv_indices,
|
dst_data_indices=dst_kv_indices,
|
||||||
executor=executor,
|
executor=executor,
|
||||||
|
# The unified pool registers ONE region holding every layer's K and
|
||||||
|
# V inside each page envelope. The MHA branch would half-split that
|
||||||
|
# single region into K and V halves and compute num_kv_layers = 0,
|
||||||
|
# transferring nothing at all; the flat branch addresses the region
|
||||||
|
# as-is. MLA-unified already reaches the flat branch via
|
||||||
|
# is_mla_backend, so this only adds the MHA-unified peer.
|
||||||
|
force_flat=get_memory().enable_unified_memory,
|
||||||
src_layer_ids=self.kv_args.kv_layer_ids,
|
src_layer_ids=self.kv_args.kv_layer_ids,
|
||||||
dst_layer_ids=dst_layer_ids,
|
dst_layer_ids=dst_layer_ids,
|
||||||
dst_device_data_indices=dst_device_kv_indices,
|
dst_device_data_indices=dst_device_kv_indices,
|
||||||
@@ -1598,8 +1605,16 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
dst_data_indices=np.array(dst_indices_local, dtype=np.int32),
|
dst_data_indices=np.array(dst_indices_local, dtype=np.int32),
|
||||||
executor=executor,
|
executor=executor,
|
||||||
state_type=st,
|
state_type=st,
|
||||||
force_flat=st
|
# Two independent reasons to keep the flat layout.
|
||||||
in (StateType.QSA_PENDING, StateType.QSA_COMPRESSED),
|
# QSA's per-layer list must not be half-split into K/V;
|
||||||
|
# neither must a unified sub-pool's single region, which
|
||||||
|
# holds every layer's K and V per slot envelope -- the
|
||||||
|
# MHA branch would compute zero layers and ship nothing
|
||||||
|
# (same reason as in `send_kvcache`).
|
||||||
|
force_flat=(
|
||||||
|
st in (StateType.QSA_PENDING, StateType.QSA_COMPRESSED)
|
||||||
|
or get_memory().enable_unified_memory
|
||||||
|
),
|
||||||
src_layer_ids=src_state_layer_ids,
|
src_layer_ids=src_state_layer_ids,
|
||||||
dst_layer_ids=dst_state_layer_ids,
|
dst_layer_ids=dst_state_layer_ids,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1328,7 +1328,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
req.kv.req_pool_idx, window_start:seq_len
|
req.kv.req_pool_idx, window_start:seq_len
|
||||||
]
|
]
|
||||||
window_kv_indices_swa = (
|
window_kv_indices_swa = (
|
||||||
self.token_to_kv_pool_allocator.translate_loc_from_full_to_swa(
|
self.token_to_kv_pool_allocator.translate_swa_indices_for_transfer(
|
||||||
window_kv_indices_full
|
window_kv_indices_full
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -346,6 +346,17 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
|
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
|
||||||
return self.logical_attn_allocator.translate_loc_from_full_to_swa(kv_indices)
|
return self.logical_attn_allocator.translate_loc_from_full_to_swa(kv_indices)
|
||||||
|
|
||||||
|
def translate_swa_indices_for_transfer(
|
||||||
|
self, kv_indices: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
# Delegated like the read-path translate above: this composite is not a
|
||||||
|
# SWA allocator itself, so it inherits neither the default nor an
|
||||||
|
# override, and the PD payload path calls this on whatever allocator
|
||||||
|
# the scheduler holds.
|
||||||
|
return self.logical_attn_allocator.translate_swa_indices_for_transfer(
|
||||||
|
kv_indices
|
||||||
|
)
|
||||||
|
|
||||||
def full_available_size(self):
|
def full_available_size(self):
|
||||||
return min(
|
return min(
|
||||||
self.logical_attn_allocator.full_available_size(),
|
self.logical_attn_allocator.full_available_size(),
|
||||||
|
|||||||
@@ -203,6 +203,19 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
assert self._kvcache.full_to_swa_index_mapping is not None
|
assert self._kvcache.full_to_swa_index_mapping is not None
|
||||||
return self._kvcache.translate_loc_from_full_to_swa(kv_indices)
|
return self._kvcache.translate_loc_from_full_to_swa(kv_indices)
|
||||||
|
|
||||||
|
def translate_swa_indices_for_transfer(
|
||||||
|
self, kv_indices: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Sliding-window token ids as the PD transfer engine addresses them.
|
||||||
|
|
||||||
|
The sibling of `translate_kv_indices_for_transfer` for the SWA state
|
||||||
|
component. On a static pool the sliding-window buffers are indexed by
|
||||||
|
the same ids the kernels use, so the read-path translate IS the answer.
|
||||||
|
A virtual-id pool must override: the transfer addresses raw bytes and
|
||||||
|
needs PHYSICAL ids, not kernel-facing ones.
|
||||||
|
"""
|
||||||
|
return self.translate_loc_from_full_to_swa(kv_indices)
|
||||||
|
|
||||||
def alloc(self, need_size: int):
|
def alloc(self, need_size: int):
|
||||||
assert self.page_size == 1
|
assert self.page_size == 1
|
||||||
if need_size > self.full_attn_allocator.available_size():
|
if need_size > self.full_attn_allocator.available_size():
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ sub-pools of one `UnifiedKVPool`, and the tri-pool variant that adds mamba state
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import List, Optional, Sequence
|
from typing import Callable, List, Optional, Sequence, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch.profiler import record_function
|
from torch.profiler import record_function
|
||||||
@@ -34,6 +34,7 @@ from sglang.srt.mem_cache.allocator.unified_sub_pool import (
|
|||||||
_flush_deferred_free_group,
|
_flush_deferred_free_group,
|
||||||
_full_tokens_before_mamba_recheck,
|
_full_tokens_before_mamba_recheck,
|
||||||
_relieve_for_alloc,
|
_relieve_for_alloc,
|
||||||
|
install_move_gate,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.unified_memory_pool import UnifiedKVPool
|
from sglang.srt.mem_cache.unified_memory_pool import UnifiedKVPool
|
||||||
from sglang.srt.utils.common import get_num_new_pages
|
from sglang.srt.utils.common import get_num_new_pages
|
||||||
@@ -320,6 +321,46 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
"""Page-level physical->virtual table of the full sub-pool."""
|
"""Page-level physical->virtual table of the full sub-pool."""
|
||||||
return self.full_attn_allocator.physical_to_virtual
|
return self.full_attn_allocator.physical_to_virtual
|
||||||
|
|
||||||
|
def translate_kv_indices_for_transfer(
|
||||||
|
self, kv_indices: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Virtual TOKEN ids -> full-sub-pool PHYSICAL token ids for the PD
|
||||||
|
transfer engine.
|
||||||
|
|
||||||
|
PHYSICAL, not kernel-facing: the transfer registers page ENVELOPES (see
|
||||||
|
`UnifiedMHATokenToKVPool.get_contiguous_buf_infos`). Without this
|
||||||
|
override the base identity would put VIRTUAL ids on the wire, which
|
||||||
|
address real bytes and so corrupt silently rather than fail.
|
||||||
|
"""
|
||||||
|
return self.full_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64))
|
||||||
|
|
||||||
|
def translate_swa_indices_for_transfer(
|
||||||
|
self, kv_indices: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Virtual TOKEN ids -> swa-sub-pool PHYSICAL token ids.
|
||||||
|
|
||||||
|
The SWA counterpart of the above. `translate_loc_from_full_to_swa`
|
||||||
|
cannot serve here: it returns KERNEL-FACING ids (the physical page
|
||||||
|
scaled by the sub-pool's per-page block count), which index the
|
||||||
|
per-layer views, whereas the SWA state component is registered as whole
|
||||||
|
page envelopes and addressed by physical page.
|
||||||
|
"""
|
||||||
|
return self.swa_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64))
|
||||||
|
|
||||||
|
def _move_gate_targets(self):
|
||||||
|
"""Every member a compaction gate must cover. A subclass that adds an
|
||||||
|
end overrides THIS, and every gate widens with it."""
|
||||||
|
return (self.full_attn_allocator, self.swa_attn_allocator)
|
||||||
|
|
||||||
|
def set_disagg_move_gate(self, gate: Callable[[], bool]) -> None:
|
||||||
|
install_move_gate(
|
||||||
|
self._move_gate_targets(),
|
||||||
|
slot="disagg_move_gate",
|
||||||
|
gate=gate,
|
||||||
|
feature="PD disaggregation",
|
||||||
|
lazy_compaction=self.lazy_compaction,
|
||||||
|
)
|
||||||
|
|
||||||
def translate_kv_loc_for_kernel(
|
def translate_kv_loc_for_kernel(
|
||||||
self,
|
self,
|
||||||
loc: torch.Tensor,
|
loc: torch.Tensor,
|
||||||
@@ -374,6 +415,52 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
self.swa_attn_allocator.alloc_with_virtual(new_virtual_pages)
|
self.swa_attn_allocator.alloc_with_virtual(new_virtual_pages)
|
||||||
return v_tokens
|
return v_tokens
|
||||||
|
|
||||||
|
def _extend_in_virtual_space(
|
||||||
|
self,
|
||||||
|
prefix_lens: torch.Tensor,
|
||||||
|
prefix_lens_cpu: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
seq_lens_cpu: torch.Tensor,
|
||||||
|
last_loc: torch.Tensor,
|
||||||
|
extend_num_tokens: int,
|
||||||
|
) -> Optional[Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
|
"""Run the full side's paged extend and report which virtual PAGES it
|
||||||
|
newly took. Returns (virtual TOKEN ids, new virtual PAGE ids), or None
|
||||||
|
when the joint capacity check cannot fund the allocation.
|
||||||
|
|
||||||
|
Both extend entries share this; they differ only in which of those pages
|
||||||
|
the sliding-window side then binds.
|
||||||
|
"""
|
||||||
|
num_new_pages = get_num_new_pages(
|
||||||
|
seq_lens=seq_lens_cpu,
|
||||||
|
page_size=self.page_size,
|
||||||
|
prefix_lens=prefix_lens_cpu,
|
||||||
|
)
|
||||||
|
need_tokens = num_new_pages * self.page_size
|
||||||
|
if need_tokens > self.available_size():
|
||||||
|
if not _relieve_for_alloc(self, need_tokens):
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Snapshot the virtual PAGES the kernel will consume; clone so swa keeps
|
||||||
|
# its view after the slice is consumed.
|
||||||
|
fa = self.full_attn_allocator
|
||||||
|
new_virtual_pages = fa.free_virtual_ids[:num_new_pages].clone()
|
||||||
|
|
||||||
|
out_indices = fa.alloc_extend(
|
||||||
|
prefix_lens,
|
||||||
|
prefix_lens_cpu,
|
||||||
|
seq_lens,
|
||||||
|
seq_lens_cpu,
|
||||||
|
last_loc,
|
||||||
|
extend_num_tokens,
|
||||||
|
num_new_pages=num_new_pages,
|
||||||
|
)
|
||||||
|
assert out_indices is not None, (
|
||||||
|
"UnifiedSWA: full.alloc_extend returned None after joint pre-check "
|
||||||
|
"passed — internal-state inconsistency"
|
||||||
|
)
|
||||||
|
return out_indices, new_virtual_pages
|
||||||
|
|
||||||
def alloc_extend(
|
def alloc_extend(
|
||||||
self,
|
self,
|
||||||
prefix_lens: torch.Tensor,
|
prefix_lens: torch.Tensor,
|
||||||
@@ -386,37 +473,75 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
"""Paged extend; returns virtual TOKEN ids. The same virtual page maps to
|
"""Paged extend; returns virtual TOKEN ids. The same virtual page maps to
|
||||||
full- and swa-physical, so swa binds exactly what the full kernel consumed."""
|
full- and swa-physical, so swa binds exactly what the full kernel consumed."""
|
||||||
with record_function("UnifiedSWAAlloc.alloc_extend"):
|
with record_function("UnifiedSWAAlloc.alloc_extend"):
|
||||||
num_new_pages = get_num_new_pages(
|
extended = self._extend_in_virtual_space(
|
||||||
seq_lens=seq_lens_cpu,
|
|
||||||
page_size=self.page_size,
|
|
||||||
prefix_lens=prefix_lens_cpu,
|
|
||||||
)
|
|
||||||
need_tokens = num_new_pages * self.page_size
|
|
||||||
if need_tokens > self.available_size():
|
|
||||||
if not _relieve_for_alloc(self, need_tokens):
|
|
||||||
return None
|
|
||||||
|
|
||||||
# Snapshot the virtual PAGES the kernel will consume; clone so swa keeps
|
|
||||||
# its view after the slice is consumed.
|
|
||||||
fa = self.full_attn_allocator
|
|
||||||
new_virtual_pages = fa.free_virtual_ids[:num_new_pages].clone()
|
|
||||||
|
|
||||||
out_indices = fa.alloc_extend(
|
|
||||||
prefix_lens,
|
prefix_lens,
|
||||||
prefix_lens_cpu,
|
prefix_lens_cpu,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
seq_lens_cpu,
|
seq_lens_cpu,
|
||||||
last_loc,
|
last_loc,
|
||||||
extend_num_tokens,
|
extend_num_tokens,
|
||||||
num_new_pages=num_new_pages,
|
|
||||||
)
|
|
||||||
assert out_indices is not None, (
|
|
||||||
"UnifiedSWA.alloc_extend: full.alloc_extend returned None "
|
|
||||||
"after joint pre-check passed — internal-state inconsistency"
|
|
||||||
)
|
)
|
||||||
|
if extended is None:
|
||||||
|
return None
|
||||||
|
out_indices, new_virtual_pages = extended
|
||||||
self.swa_attn_allocator.alloc_with_virtual(new_virtual_pages)
|
self.swa_attn_allocator.alloc_with_virtual(new_virtual_pages)
|
||||||
return out_indices # virtual TOKEN ids
|
return out_indices # virtual TOKEN ids
|
||||||
|
|
||||||
|
def alloc_extend_swa_tail(
|
||||||
|
self,
|
||||||
|
prefix_lens: torch.Tensor,
|
||||||
|
prefix_lens_cpu: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
seq_lens_cpu: torch.Tensor,
|
||||||
|
last_loc: torch.Tensor,
|
||||||
|
extend_num_tokens: int,
|
||||||
|
swa_tail_len: int,
|
||||||
|
) -> Optional[torch.Tensor]:
|
||||||
|
"""Decode-node prealloc: full KV for the whole sequence, sliding-window
|
||||||
|
KV for the live window tail only.
|
||||||
|
|
||||||
|
The static composite allocates the two sides independently and records
|
||||||
|
a full->swa index mapping. That is not representable here: the two
|
||||||
|
sides SHARE one virtual id space (a virtual page names a full-physical
|
||||||
|
page and, if bound, a swa-physical one), which is why
|
||||||
|
`set_full_to_swa_mapping` is a no-op on this allocator and
|
||||||
|
`translate_loc_from_full_to_swa` derives the swa id from the virtual id
|
||||||
|
instead of a table. Running the static body would call `alloc_extend`
|
||||||
|
on the swa sub-allocator, which asserts it is not the id owner.
|
||||||
|
|
||||||
|
The tail is expressed by binding swa for the TAIL's virtual pages only.
|
||||||
|
A new page left unbound has no swa-physical page, which reads as the
|
||||||
|
sink and is skipped by `free`'s `swa_v2p_page > 0` mask -- exactly the
|
||||||
|
out-of-window state the ratchet produces via `free_swa`.
|
||||||
|
|
||||||
|
Admission is priced at the FULL side's page count, as plain
|
||||||
|
`alloc_extend` is: pessimistic when the tail is short, but it reuses
|
||||||
|
the composite's audited joint capacity path, and the bytes actually
|
||||||
|
held still follow the tail.
|
||||||
|
"""
|
||||||
|
assert len(prefix_lens_cpu) == 1
|
||||||
|
assert 0 <= swa_tail_len <= extend_num_tokens
|
||||||
|
with record_function("UnifiedSWAAlloc.alloc_extend_swa_tail"):
|
||||||
|
extended = self._extend_in_virtual_space(
|
||||||
|
prefix_lens,
|
||||||
|
prefix_lens_cpu,
|
||||||
|
seq_lens,
|
||||||
|
seq_lens_cpu,
|
||||||
|
last_loc,
|
||||||
|
extend_num_tokens,
|
||||||
|
)
|
||||||
|
if extended is None:
|
||||||
|
return None
|
||||||
|
out_indices, new_virtual_pages = extended
|
||||||
|
if swa_tail_len > 0 and new_virtual_pages.numel() > 0:
|
||||||
|
tail_pages = torch.unique(out_indices[-swa_tail_len:] // self.page_size)
|
||||||
|
# Only NEW pages need binding; a tail page carried in from the
|
||||||
|
# prefix is already bound on the swa side.
|
||||||
|
to_bind = new_virtual_pages[torch.isin(new_virtual_pages, tail_pages)]
|
||||||
|
if to_bind.numel() > 0:
|
||||||
|
self.swa_attn_allocator.alloc_with_virtual(to_bind)
|
||||||
|
return out_indices # virtual TOKEN ids
|
||||||
|
|
||||||
def alloc_decode(
|
def alloc_decode(
|
||||||
self,
|
self,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
@@ -769,16 +894,9 @@ class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWATokenToKVPoolAllocator):
|
|||||||
binary search; the order matches the alloc path (full takes the high band).
|
binary search; the order matches the alloc path (full takes the high band).
|
||||||
"""
|
"""
|
||||||
fa, sa = self.full_attn_allocator, self.swa_attn_allocator
|
fa, sa = self.full_attn_allocator, self.swa_attn_allocator
|
||||||
e_f, e_s = fa.entry_bytes_per_page, sa.entry_bytes_per_page
|
e_f = fa.entry_bytes_per_page
|
||||||
# full is grow-down: its chain gap IS the high band.
|
# full is grow-down: its chain gap IS the high band.
|
||||||
b_high = fa._current_gap_bytes()
|
b_high = fa._current_gap_bytes()
|
||||||
if sa._is_frontier_transparent():
|
|
||||||
b_low = 0
|
|
||||||
else:
|
|
||||||
b_low = max(
|
|
||||||
0,
|
|
||||||
sa._byte_low_frontier() - sa._chain_high_frontier_below_bytes(),
|
|
||||||
)
|
|
||||||
h_f = len(fa._free_phys_pages) if fa.lazy_compaction else 0
|
h_f = len(fa._free_phys_pages) if fa.lazy_compaction else 0
|
||||||
h_s = sa._hole_pages()
|
h_s = sa._hole_pages()
|
||||||
r_f = fa.num_pages - fa.min_page_index - fa._allocated_pages()
|
r_f = fa.num_pages - fa.min_page_index - fa._allocated_pages()
|
||||||
@@ -819,6 +937,16 @@ class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWATokenToKVPoolAllocator):
|
|||||||
hi_n = mid - 1
|
hi_n = mid - 1
|
||||||
return lo_n * self.page_size
|
return lo_n * self.page_size
|
||||||
|
|
||||||
|
def _move_gate_targets(self):
|
||||||
|
"""All three members. The mamba end compacts independently and its slot
|
||||||
|
envelopes move as `StateType.MAMBA`, so leaving it out of a gate would
|
||||||
|
let a conv/SSM slot relocate under an in-flight transfer."""
|
||||||
|
return (
|
||||||
|
self.full_attn_allocator,
|
||||||
|
self.swa_attn_allocator,
|
||||||
|
self.mamba_allocator,
|
||||||
|
)
|
||||||
|
|
||||||
def _flush_targets(self):
|
def _flush_targets(self):
|
||||||
"""All three members, float FIRST: its zero-copy boundary absorption must
|
"""All three members, float FIRST: its zero-copy boundary absorption must
|
||||||
land before the deficit math prices a relocation it already covered."""
|
land before the deficit math prices a relocation it already covered."""
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from sglang.srt.mem_cache.allocator.unified_sub_pool import (
|
|||||||
_flush_deferred_free_group,
|
_flush_deferred_free_group,
|
||||||
_full_tokens_before_mamba_recheck,
|
_full_tokens_before_mamba_recheck,
|
||||||
_relieve_for_alloc,
|
_relieve_for_alloc,
|
||||||
|
install_move_gate,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.unified_memory_pool import UnifiedKVPool
|
from sglang.srt.mem_cache.unified_memory_pool import UnifiedKVPool
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
@@ -306,15 +307,20 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
)
|
)
|
||||||
return self.full_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64))
|
return self.full_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64))
|
||||||
|
|
||||||
|
def _move_gate_targets(self):
|
||||||
|
"""Every member a compaction gate must cover. The mamba end is gated
|
||||||
|
even where its state is not itself transferred: the gate is about the
|
||||||
|
MOVER, and the two ends compact as peers."""
|
||||||
|
return (self.full_attn_allocator, self.mamba_allocator)
|
||||||
|
|
||||||
def set_disagg_move_gate(self, gate: Callable[[], bool]) -> None:
|
def set_disagg_move_gate(self, gate: Callable[[], bool]) -> None:
|
||||||
"""Install the PD-disaggregation move gate on both sub-allocators."""
|
install_move_gate(
|
||||||
assert self.lazy_compaction, (
|
self._move_gate_targets(),
|
||||||
"PD disaggregation with the unified memory pool requires lazy "
|
slot="disagg_move_gate",
|
||||||
"compaction (eager free-path compaction moves pages under "
|
gate=gate,
|
||||||
"in-flight transfers)."
|
feature="PD disaggregation",
|
||||||
|
lazy_compaction=self.lazy_compaction,
|
||||||
)
|
)
|
||||||
self.full_attn_allocator.disagg_move_gate = gate
|
|
||||||
self.mamba_allocator.disagg_move_gate = gate
|
|
||||||
|
|
||||||
def is_slot_allocated(self, slot: int) -> bool:
|
def is_slot_allocated(self, slot: int) -> bool:
|
||||||
return self.full_attn_allocator.is_slot_allocated(slot)
|
return self.full_attn_allocator.is_slot_allocated(slot)
|
||||||
|
|||||||
@@ -239,6 +239,31 @@ def _full_tokens_before_mamba_recheck(
|
|||||||
return -(-minimum_missing_bytes * dcp_size // full_allocator.entry_bytes)
|
return -(-minimum_missing_bytes * dcp_size // full_allocator.entry_bytes)
|
||||||
|
|
||||||
|
|
||||||
|
def install_move_gate(
|
||||||
|
targets,
|
||||||
|
*,
|
||||||
|
slot: str,
|
||||||
|
gate: Callable[[], bool],
|
||||||
|
feature: str,
|
||||||
|
lazy_compaction: bool,
|
||||||
|
) -> None:
|
||||||
|
"""Point every member of a composite at one compaction gate.
|
||||||
|
|
||||||
|
A gate that reaches only some members is not a weaker gate, it is no gate:
|
||||||
|
the ungated end relocates its own pages under the same in-flight transfer.
|
||||||
|
So the member list is stated once per composite (`_move_gate_targets`) and
|
||||||
|
every gate installs over it, rather than each setter naming the members it
|
||||||
|
happens to remember.
|
||||||
|
"""
|
||||||
|
assert lazy_compaction, (
|
||||||
|
f"{feature} with the unified memory pool requires lazy compaction "
|
||||||
|
"(eager free-path compaction moves pages under in-flight transfers)."
|
||||||
|
)
|
||||||
|
assert slot in ("disagg_move_gate", "host_transfer_move_gate"), slot
|
||||||
|
for target in targets:
|
||||||
|
setattr(target, slot, gate)
|
||||||
|
|
||||||
|
|
||||||
class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||||
"""Allocator for one sub-pool over a `UnifiedKVPool`."""
|
"""Allocator for one sub-pool over a `UnifiedKVPool`."""
|
||||||
|
|
||||||
|
|||||||
@@ -155,11 +155,21 @@ def resolve_decode_retraction_backup(*, tp_worker: BaseTpWorker) -> str:
|
|||||||
)
|
)
|
||||||
# Host-pool retraction transfers full and sliding-window components
|
# Host-pool retraction transfers full and sliding-window components
|
||||||
# only, so a model with recurrent state stays on cpu_tensor.
|
# only, so a model with recurrent state stays on cpu_tensor.
|
||||||
supports_host_pool = not uses_ssm_state(
|
#
|
||||||
tp_worker.model_runner.model_config
|
# The unified pool is excluded for the same reason hierarchical cache is
|
||||||
) and (
|
# (see `handle_unified_memory_pool`): the host-transfer path indexes the
|
||||||
isinstance(kv_cache, MHATokenToKVPool)
|
# device buffers with the ids it is handed, and under the unified pool
|
||||||
or (isinstance(kv_cache, SWAKVPool) and full_tokens_per_layer > 0)
|
# those are VIRTUAL. It also cannot be sized from `kv_cache.size`, which
|
||||||
|
# is a KERNEL-FACING row count (`num_pages * 2 * layer_num * page_size`)
|
||||||
|
# rather than a token capacity -- gpt-oss-20b reports 85M "tokens" and
|
||||||
|
# asks for 418 GB of host memory per component.
|
||||||
|
supports_host_pool = (
|
||||||
|
not uses_ssm_state(tp_worker.model_runner.model_config)
|
||||||
|
and not memory.enable_unified_memory
|
||||||
|
and (
|
||||||
|
isinstance(kv_cache, MHATokenToKVPool)
|
||||||
|
or (isinstance(kv_cache, SWAKVPool) and full_tokens_per_layer > 0)
|
||||||
|
)
|
||||||
)
|
)
|
||||||
schedule = get_schedule()
|
schedule = get_schedule()
|
||||||
priority_preemption = (
|
priority_preemption = (
|
||||||
|
|||||||
@@ -446,24 +446,12 @@ class KVCacheConfigurator:
|
|||||||
# from one byte buffer, then return. Gated to the target worker
|
# from one byte buffer, then return. Gated to the target worker
|
||||||
# (req_to_token_pool is None); supports hybrid Mamba and hybrid SWA (not DSV4).
|
# (req_to_token_pool is None); supports hybrid Mamba and hybrid SWA (not DSV4).
|
||||||
if get_memory().enable_unified_memory and req_to_token_pool is None:
|
if get_memory().enable_unified_memory and req_to_token_pool is None:
|
||||||
pd_enabled = get_disagg().disaggregation_mode != "null"
|
|
||||||
is_dsv4 = is_deepseek_v4(self.model_config.hf_config)
|
is_dsv4 = is_deepseek_v4(self.model_config.hf_config)
|
||||||
# Order matters: an Inkling-class model is BOTH mambaish and
|
# Order matters: an Inkling-class model is BOTH mambaish and
|
||||||
# hybrid-SWA, and the mamba pair would store every SWA layer's KV at
|
# hybrid-SWA, and the mamba pair would store every SWA layer's KV at
|
||||||
# FULL lifetime -- its branch reads the HF config's
|
# FULL lifetime -- its branch reads the HF config's
|
||||||
# full_attention_layer_ids, which for Inkling is ALL layers.
|
# full_attention_layer_ids, which for Inkling is ALL layers.
|
||||||
if self.mambaish_config is not None and self.is_hybrid_swa and not is_dsv4:
|
if self.mambaish_config is not None and self.is_hybrid_swa and not is_dsv4:
|
||||||
if pd_enabled:
|
|
||||||
# Same limitation as the 2-pool SWA branch below: the
|
|
||||||
# tri-pool carries an SWA sub-pool, and there is no
|
|
||||||
# whole-envelope transfer scheme for it.
|
|
||||||
raise ValueError(
|
|
||||||
"--enable-unified-memory with PD disaggregation does "
|
|
||||||
"not support hybrid-SWA models yet (no whole-envelope "
|
|
||||||
"transfer scheme for the SWA sub-pool); this model "
|
|
||||||
"routes to the mamba+SWA tri-pool, which has one. Drop "
|
|
||||||
"--enable-unified-memory or run without PD."
|
|
||||||
)
|
|
||||||
bundle = self._init_unified_mamba_swa_pools(
|
bundle = self._init_unified_mamba_swa_pools(
|
||||||
max_num_reqs=sizes.max_running_requests,
|
max_num_reqs=sizes.max_running_requests,
|
||||||
full_max_total_num_tokens=sizes.full_max_total_num_tokens,
|
full_max_total_num_tokens=sizes.full_max_total_num_tokens,
|
||||||
@@ -471,27 +459,12 @@ class KVCacheConfigurator:
|
|||||||
unified_total_bytes=sizes.unified_total_bytes,
|
unified_total_bytes=sizes.unified_total_bytes,
|
||||||
)
|
)
|
||||||
elif self.mambaish_config is not None:
|
elif self.mambaish_config is not None:
|
||||||
if pd_enabled and not self.use_mla_backend:
|
|
||||||
raise ValueError(
|
|
||||||
"--enable-unified-memory with PD disaggregation "
|
|
||||||
"currently supports only MLA hybrid-Mamba models "
|
|
||||||
"(e.g. kimi-linear); this model uses the MHA full-"
|
|
||||||
"attention pool. Drop --enable-unified-memory or run "
|
|
||||||
"without PD disaggregation."
|
|
||||||
)
|
|
||||||
bundle = self._init_unified_mamba_pools(
|
bundle = self._init_unified_mamba_pools(
|
||||||
max_num_reqs=sizes.max_running_requests,
|
max_num_reqs=sizes.max_running_requests,
|
||||||
max_total_num_tokens=sizes.max_total_num_tokens,
|
max_total_num_tokens=sizes.max_total_num_tokens,
|
||||||
unified_total_bytes=sizes.unified_total_bytes,
|
unified_total_bytes=sizes.unified_total_bytes,
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_swa and not is_dsv4:
|
elif self.is_hybrid_swa and not is_dsv4:
|
||||||
if pd_enabled:
|
|
||||||
raise ValueError(
|
|
||||||
"--enable-unified-memory with PD disaggregation does "
|
|
||||||
"not support hybrid-SWA models yet (no whole-envelope "
|
|
||||||
"transfer scheme for the SWA sub-pool). Drop "
|
|
||||||
"--enable-unified-memory or run without PD."
|
|
||||||
)
|
|
||||||
bundle = self._init_unified_swa_pools(
|
bundle = self._init_unified_swa_pools(
|
||||||
max_num_reqs=sizes.max_running_requests,
|
max_num_reqs=sizes.max_running_requests,
|
||||||
full_max_total_num_tokens=sizes.full_max_total_num_tokens,
|
full_max_total_num_tokens=sizes.full_max_total_num_tokens,
|
||||||
@@ -835,6 +808,13 @@ class KVCacheConfigurator:
|
|||||||
unified_total_bytes=(None if self.is_draft_worker else unified_total_bytes),
|
unified_total_bytes=(None if self.is_draft_worker else unified_total_bytes),
|
||||||
# bs=1 feasibility floor input (context len is already passed).
|
# bs=1 feasibility floor input (context len is already passed).
|
||||||
sliding_window_size=self.model_config.sliding_window_size,
|
sliding_window_size=self.model_config.sliding_window_size,
|
||||||
|
# Decode nodes hand out request rows to PREALLOCATED transfers on
|
||||||
|
# top of the running set; the 2-pool mamba factory takes the same.
|
||||||
|
decode_pre_alloc_size=(
|
||||||
|
get_disagg().disaggregation_decode_extra_slots
|
||||||
|
if get_disagg().disaggregation_mode == "decode"
|
||||||
|
else 0
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_unified_swa_pools(
|
def _init_unified_swa_pools(
|
||||||
@@ -863,12 +843,28 @@ class KVCacheConfigurator:
|
|||||||
extra_max_context_len = 4
|
extra_max_context_len = 4
|
||||||
if get_spec().speculative_num_draft_tokens is not None:
|
if get_spec().speculative_num_draft_tokens is not None:
|
||||||
extra_max_context_len += get_spec().speculative_num_draft_tokens
|
extra_max_context_len += get_spec().speculative_num_draft_tokens
|
||||||
req_to_token_pool = ReqToTokenPool(
|
if get_disagg().disaggregation_mode == "decode":
|
||||||
size=max_num_reqs,
|
# A decode node hands out request rows to PREALLOCATED transfers on
|
||||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
# top of its running set, so it needs the extra-slot pool (and the
|
||||||
device=self.device,
|
# `pre_alloc_size` the scheduler's invariant checker reads). Mirrors
|
||||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
# `_build_req_to_token_pool`'s decode branch; the mamba composite
|
||||||
)
|
# already takes `decode_pre_alloc_size` the same way.
|
||||||
|
from sglang.srt.disaggregation.decode import DecodeReqToTokenPool
|
||||||
|
|
||||||
|
req_to_token_pool = DecodeReqToTokenPool(
|
||||||
|
size=max_num_reqs,
|
||||||
|
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||||
|
device=self.device,
|
||||||
|
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||||
|
pre_alloc_size=get_disagg().disaggregation_decode_extra_slots,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
req_to_token_pool = ReqToTokenPool(
|
||||||
|
size=max_num_reqs,
|
||||||
|
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||||
|
device=self.device,
|
||||||
|
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||||
|
)
|
||||||
|
|
||||||
head_num = self.model_config.get_num_kv_heads(
|
head_num = self.model_config.get_num_kv_heads(
|
||||||
get_parallel().attn_tp_size, get_parallel().attn_dcp_size
|
get_parallel().attn_tp_size, get_parallel().attn_dcp_size
|
||||||
|
|||||||
@@ -624,10 +624,19 @@ class UnifiedMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
env[tgt_pages] = env[src_pages]
|
env[tgt_pages] = env[src_pages]
|
||||||
|
|
||||||
def get_contiguous_buf_infos(self):
|
def get_contiguous_buf_infos(self):
|
||||||
raise NotImplementedError(
|
"""PD-transfer registration: ONE entry, the raw buffer, addressed as
|
||||||
"unified layout has no per-layer contiguous regions; "
|
``raw_ptr + physical_page_id * page_envelope_bytes``.
|
||||||
"KV transfer / disaggregation is unsupported."
|
|
||||||
)
|
Same whole-envelope contract as `UnifiedMLATokenToKVPool`: the transfer
|
||||||
|
item is one page across ALL layers and both K and V, because the
|
||||||
|
per-layer views overlap inside the envelope and index in kernel-facing
|
||||||
|
ids. A peer must therefore build an identical spec -- enforced on the
|
||||||
|
wire by `_validate_envelope_kv_layout`.
|
||||||
|
"""
|
||||||
|
# The address formula omits the anchor; a nonzero one would mis-address.
|
||||||
|
assert self._unified_buffer.anchor_bytes(self._sub_pool_name) == 0
|
||||||
|
raw = self._unified_buffer._raw
|
||||||
|
return [raw.data_ptr()], [raw.numel()], [self._page_bytes]
|
||||||
|
|
||||||
def get_cpu_copy(self, indices, mamba_indices=None):
|
def get_cpu_copy(self, indices, mamba_indices=None):
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
@@ -1882,6 +1891,7 @@ def init_unified_mamba_swa_pools(
|
|||||||
lazy_compaction: bool = False,
|
lazy_compaction: bool = False,
|
||||||
unified_total_bytes: Optional[int] = None,
|
unified_total_bytes: Optional[int] = None,
|
||||||
sliding_window_size: Optional[int] = None,
|
sliding_window_size: Optional[int] = None,
|
||||||
|
decode_pre_alloc_size: int = 0,
|
||||||
) -> UnifiedPoolBundle:
|
) -> UnifiedPoolBundle:
|
||||||
"""Build the TRI-pool unified-memory-pool stack for models with full KV +
|
"""Build the TRI-pool unified-memory-pool stack for models with full KV +
|
||||||
SWA KV + mamba/conv state (Inkling-class: `mambaish_config` AND
|
SWA KV + mamba/conv state (Inkling-class: `mambaish_config` AND
|
||||||
@@ -2007,6 +2017,7 @@ def init_unified_mamba_swa_pools(
|
|||||||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||||||
enable_overlap_schedule=not disable_overlap_schedule,
|
enable_overlap_schedule=not disable_overlap_schedule,
|
||||||
start_layer=start_layer,
|
start_layer=start_layer,
|
||||||
|
pre_alloc_size=decode_pre_alloc_size,
|
||||||
)
|
)
|
||||||
allocator = UnifiedMambaSWATokenToKVPoolAllocator(
|
allocator = UnifiedMambaSWATokenToKVPoolAllocator(
|
||||||
unified_buffer=shared_pool,
|
unified_buffer=shared_pool,
|
||||||
|
|||||||
@@ -9,6 +9,9 @@ from sglang.test.server_fixtures.disaggregation_fixture import (
|
|||||||
register_cuda_ci(est_time=236, stage="base-b", runner_config="2-gpu-large")
|
register_cuda_ci(est_time=236, stage="base-b", runner_config="2-gpu-large")
|
||||||
|
|
||||||
KIMI_LINEAR_MODEL = "yujiepan/kimi-linear-tiny-random"
|
KIMI_LINEAR_MODEL = "yujiepan/kimi-linear-tiny-random"
|
||||||
|
# Smallest in-tree GDN hybrid: MHA full attention + gated-delta-net linear
|
||||||
|
# layers, i.e. the unified pool's MHA sub-pool rather than the MLA one.
|
||||||
|
QWEN_GDN_MODEL = "Qwen/Qwen3.5-0.8B"
|
||||||
SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"}
|
SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"}
|
||||||
|
|
||||||
# --attention-backend and --enable-deterministic-inference are deliberately
|
# --attention-backend and --enable-deterministic-inference are deliberately
|
||||||
@@ -63,5 +66,21 @@ class TestUnifiedMemoryDisaggregationChunkedPrefill(TestUnifiedMemoryDisaggregat
|
|||||||
extra_decode_args = _chunked_args
|
extra_decode_args = _chunked_args
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedMemoryDisaggregationMHA(TestUnifiedMemoryDisaggregation):
|
||||||
|
"""The MHA full-attention sub-pool over the wire.
|
||||||
|
|
||||||
|
Kimi-Linear above exercises the MLA sub-pool, whose whole-envelope
|
||||||
|
registration has always been the one PD supports. An MHA envelope is a
|
||||||
|
different shape -- `2 * layer_num` row-blocks per page instead of
|
||||||
|
`layer_num` -- and it reaches a different branch of
|
||||||
|
`_send_kvcache_generic`: without `force_flat` the MHA branch halves the
|
||||||
|
single registered region into K and V, computes `num_kv_layers = 0` and
|
||||||
|
transfers NOTHING, which shows up as garbage decode rather than an error.
|
||||||
|
Logprob parity against a non-PD unified reference is what catches that.
|
||||||
|
"""
|
||||||
|
|
||||||
|
model = QWEN_GDN_MODEL
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,67 @@
|
|||||||
|
"""PD disaggregation for a hybrid-SWA model on the unified memory pool.
|
||||||
|
|
||||||
|
A hybrid-SWA model ships TWO attention components: the full-attention KV on the
|
||||||
|
ordinary `kv_data_ptrs` channel and the sliding-window KV as `StateType.SWA`.
|
||||||
|
Under `--enable-unified-memory` both are whole page envelopes into the SAME raw
|
||||||
|
buffer, distinguished only by their per-page stride, and each is addressed by
|
||||||
|
its OWN sub-pool's physical page id -- the full and SWA sides run independent
|
||||||
|
compactions, so one virtual token names two unrelated physical pages.
|
||||||
|
|
||||||
|
That makes three ways to be silently wrong rather than loud:
|
||||||
|
* shipping virtual ids (the base `translate_kv_indices_for_transfer` is the
|
||||||
|
identity, and virtual ids address real bytes);
|
||||||
|
* shipping the SWA side's KERNEL-FACING ids, which the read path uses, in
|
||||||
|
place of its physical ones;
|
||||||
|
* letting compaction relocate a page mid-transfer, which the SWA allocator
|
||||||
|
had no `set_disagg_move_gate` to prevent.
|
||||||
|
|
||||||
|
Logprob parity against a non-PD unified reference catches all three; GSM8K on
|
||||||
|
gpt-oss is too noisy to (single-server unified and static both score 0.570 at
|
||||||
|
200 questions, and PD runs of each span 0.540-0.610).
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.kits.pd_parity_kit import PDLogprobParityMixin
|
||||||
|
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||||
|
PDDisaggregationServerBase,
|
||||||
|
)
|
||||||
|
from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST_MXFP4_WITH_MOE
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=1200, stage="extra-a", runner_config="2-gpu-large")
|
||||||
|
|
||||||
|
UNIFIED_SWA_ARGS = [
|
||||||
|
"--skip-tokenizer-init",
|
||||||
|
"--random-seed",
|
||||||
|
"1",
|
||||||
|
"--enable-unified-memory",
|
||||||
|
# gpt-oss uses attention sinks, which flashinfer does not support; triton
|
||||||
|
# reads both sub-pools' per-layer views.
|
||||||
|
"--attention-backend",
|
||||||
|
"triton",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
|
"--cuda-graph-backend-decode",
|
||||||
|
"disabled",
|
||||||
|
"--cuda-graph-backend-prefill",
|
||||||
|
"disabled",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedMemoryDisaggregationSWA(
|
||||||
|
PDLogprobParityMixin, PDDisaggregationServerBase
|
||||||
|
):
|
||||||
|
"""1 prefill + 1 decode, both unified, vs a non-PD unified reference."""
|
||||||
|
|
||||||
|
model = DEFAULT_MODEL_NAME_FOR_TEST_MXFP4_WITH_MOE
|
||||||
|
prefill_tp_size = 1
|
||||||
|
decode_tp_size = 1
|
||||||
|
decode_base_gpu_id = 1
|
||||||
|
baseline_args = UNIFIED_SWA_ARGS
|
||||||
|
extra_prefill_args = UNIFIED_SWA_ARGS
|
||||||
|
extra_decode_args = UNIFIED_SWA_ARGS
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,86 @@
|
|||||||
|
"""PD disaggregation for a TRI-pool model on the unified memory pool.
|
||||||
|
|
||||||
|
Inkling is the only in-tree architecture that is both mambaish and hybrid-SWA,
|
||||||
|
so one unified buffer carries three components with three independent
|
||||||
|
compactions -- ``[conv state (up END) | swa (FLOAT) | full (down END)]`` -- and
|
||||||
|
PD must ship all three per request: full KV on the ``kv_data_ptrs`` channel,
|
||||||
|
sliding-window KV as ``StateType.SWA``, ShortConv state as ``StateType.MAMBA``
|
||||||
|
(via the ``req_to_token_pool`` fallback, since the KV pool here is a
|
||||||
|
``UnifiedSWAKVPool`` rather than a ``HybridLinearKVPool``).
|
||||||
|
|
||||||
|
Two failures this pins that the 2-pool cases cannot:
|
||||||
|
|
||||||
|
* the FLOAT sub-pool moves for reasons neither END does, so a move gate that
|
||||||
|
reaches only full and swa still lets a conv slot relocate under an
|
||||||
|
in-flight state transfer;
|
||||||
|
* ``page_size > 1`` turns on the decode node's SWA-tail prealloc, whose
|
||||||
|
static body allocates the swa side independently -- an assertion failure
|
||||||
|
against this composite's single virtual id space, and, once that is
|
||||||
|
handled, the first path that can bind the WRONG swa pages.
|
||||||
|
|
||||||
|
Logprob parity against a non-PD unified reference is the check: the tiny
|
||||||
|
``test`` revision is undertrained, so answer quality carries no signal, but a
|
||||||
|
dropped or misaddressed component moves logprobs immediately.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.kits.pd_parity_kit import PDLogprobParityMixin
|
||||||
|
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||||
|
PDDisaggregationServerBase,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=900, stage="extra-a", runner_config="2-gpu-large")
|
||||||
|
|
||||||
|
_MODEL_PATH = os.environ.get("INKLING_TEST_MODEL_PATH", "thinkingmachines/Inkling")
|
||||||
|
_MODEL_REVISION = os.environ.get("INKLING_TEST_MODEL_REVISION", "test")
|
||||||
|
|
||||||
|
# The unified radix tree is what merges the three components into one tree.
|
||||||
|
SERVER_ENV = {"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"}
|
||||||
|
|
||||||
|
UNIFIED_TRI_ARGS = [
|
||||||
|
"--skip-tokenizer-init",
|
||||||
|
"--random-seed",
|
||||||
|
"1",
|
||||||
|
"--enable-unified-memory",
|
||||||
|
# Unified requires the Triton strided page-major read/write paths.
|
||||||
|
"--attention-backend",
|
||||||
|
"triton",
|
||||||
|
"--page-size",
|
||||||
|
"128",
|
||||||
|
"--mamba-radix-cache-strategy",
|
||||||
|
"extra_buffer",
|
||||||
|
"--swa-full-tokens-ratio",
|
||||||
|
"0.1",
|
||||||
|
"--mamba-full-memory-ratio",
|
||||||
|
"0.1",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.5",
|
||||||
|
# Inkling defaults to a FULL prefill graph, which unified rejects at boot.
|
||||||
|
"--cuda-graph-backend-prefill",
|
||||||
|
"disabled",
|
||||||
|
"--revision",
|
||||||
|
_MODEL_REVISION,
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedMemoryDisaggregationTriPool(
|
||||||
|
PDLogprobParityMixin, PDDisaggregationServerBase
|
||||||
|
):
|
||||||
|
"""1 prefill + 1 decode, both unified, vs a non-PD unified reference."""
|
||||||
|
|
||||||
|
model = _MODEL_PATH
|
||||||
|
extra_prefill_env = SERVER_ENV
|
||||||
|
extra_decode_env = SERVER_ENV
|
||||||
|
prefill_tp_size = 1
|
||||||
|
decode_tp_size = 1
|
||||||
|
decode_base_gpu_id = 1
|
||||||
|
baseline_args = UNIFIED_TRI_ARGS
|
||||||
|
extra_prefill_args = UNIFIED_TRI_ARGS
|
||||||
|
extra_decode_args = UNIFIED_TRI_ARGS
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -186,5 +186,210 @@ class TestMoveGateRejectsNonPdNode(CustomTestCase):
|
|||||||
unified_memory_disagg_move_gate(scheduler)
|
unified_memory_disagg_move_gate(scheduler)
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedAllocatorsPublishTheTransferContract(CustomTestCase):
|
||||||
|
"""Every unified composite allocator must OVERRIDE the two PD hooks.
|
||||||
|
|
||||||
|
`BaseTokenToKVPoolAllocator.translate_kv_indices_for_transfer` is the
|
||||||
|
IDENTITY, and `set_disagg_move_gate` exists only where a composite defines
|
||||||
|
it. Inheriting either is silent, not loud: identity puts VIRTUAL ids on the
|
||||||
|
wire (they address real bytes, so the peer gets plausible garbage), and a
|
||||||
|
missing gate lets lazy compaction relocate pages under in-flight RDMA.
|
||||||
|
An AST-level check because instantiating these composites needs a GPU.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Composites that own the full-side virtual ids and so must define the
|
||||||
|
# transfer translate themselves.
|
||||||
|
_COMPOSITES = (
|
||||||
|
"UnifiedMambaTokenToKVPoolAllocator",
|
||||||
|
"UnifiedSWATokenToKVPoolAllocator",
|
||||||
|
)
|
||||||
|
# Every composite must define the gate setter, including the tri-pool,
|
||||||
|
# which inherits the SWA translates (same full side) but has a THIRD
|
||||||
|
# member the 2-pool setter does not reach.
|
||||||
|
_GATE_COMPOSITES = _COMPOSITES + ("UnifiedMambaSWATokenToKVPoolAllocator",)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _own_methods(cls_name: str) -> Set[str]:
|
||||||
|
"""Names this class defines ITSELF, inheritance excluded.
|
||||||
|
|
||||||
|
Resolved off the class object rather than by parsing a named module:
|
||||||
|
these composites have already been moved once (out of
|
||||||
|
`multi_ended_allocator` into `allocator/unified_*`), and a hardcoded
|
||||||
|
module path turns that kind of move into a test failure that says
|
||||||
|
nothing about the contract. `__dict__` needs no GPU -- it is the class
|
||||||
|
body, not an instance.
|
||||||
|
"""
|
||||||
|
from sglang.srt.mem_cache.allocator import (
|
||||||
|
unified_hybrid_swa,
|
||||||
|
unified_mamba,
|
||||||
|
)
|
||||||
|
|
||||||
|
for mod in (unified_mamba, unified_hybrid_swa):
|
||||||
|
cls = getattr(mod, cls_name, None)
|
||||||
|
if cls is not None:
|
||||||
|
return set(vars(cls))
|
||||||
|
raise AssertionError(f"class {cls_name} not found in the unified allocators")
|
||||||
|
|
||||||
|
def test_transfer_translate_is_not_inherited_identity(self):
|
||||||
|
for name in self._COMPOSITES:
|
||||||
|
with self.subTest(composite=name):
|
||||||
|
self.assertIn(
|
||||||
|
"translate_kv_indices_for_transfer",
|
||||||
|
self._own_methods(name),
|
||||||
|
f"{name} inherits the identity transfer translate; PD would "
|
||||||
|
"ship VIRTUAL ids and corrupt KV without any error",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Every sub-allocator attribute a composite can hold. The stub carries all
|
||||||
|
# of them regardless of composite, so the assertion is on what installation
|
||||||
|
# REACHES rather than on what the stub was given.
|
||||||
|
_MEMBER_ATTRS = ("full_attn_allocator", "swa_attn_allocator", "mamba_allocator")
|
||||||
|
|
||||||
|
# The members each composite's gate must reach. The tri-pool row is the one
|
||||||
|
# that matters: it inherits the setter, so an enumeration written inside
|
||||||
|
# that setter would silently leave the third member ungated.
|
||||||
|
_EXPECTED_COVERAGE = {
|
||||||
|
"UnifiedMambaTokenToKVPoolAllocator": {
|
||||||
|
"full_attn_allocator",
|
||||||
|
"mamba_allocator",
|
||||||
|
},
|
||||||
|
"UnifiedSWATokenToKVPoolAllocator": {
|
||||||
|
"full_attn_allocator",
|
||||||
|
"swa_attn_allocator",
|
||||||
|
},
|
||||||
|
"UnifiedMambaSWATokenToKVPoolAllocator": {
|
||||||
|
"full_attn_allocator",
|
||||||
|
"swa_attn_allocator",
|
||||||
|
"mamba_allocator",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
def _members_reached(self, cls_name: str, slot: str) -> Set[str]:
|
||||||
|
"""Install one gate on a stub composite and report which members got it.
|
||||||
|
|
||||||
|
`object.__new__` skips `__init__` (which needs a GPU); the setter reads
|
||||||
|
only `lazy_compaction` and the member attributes.
|
||||||
|
"""
|
||||||
|
from sglang.srt.mem_cache.allocator import unified_hybrid_swa, unified_mamba
|
||||||
|
|
||||||
|
cls = getattr(unified_mamba, cls_name, None) or getattr(
|
||||||
|
unified_hybrid_swa, cls_name
|
||||||
|
)
|
||||||
|
alloc = object.__new__(cls)
|
||||||
|
alloc.lazy_compaction = True
|
||||||
|
for attr in self._MEMBER_ATTRS:
|
||||||
|
member = type("_Member", (), {})()
|
||||||
|
member.disagg_move_gate = None
|
||||||
|
member.host_transfer_move_gate = None
|
||||||
|
setattr(alloc, attr, member)
|
||||||
|
|
||||||
|
def gate() -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
alloc.set_disagg_move_gate(gate)
|
||||||
|
return {
|
||||||
|
attr
|
||||||
|
for attr in self._MEMBER_ATTRS
|
||||||
|
if getattr(getattr(alloc, attr), slot) is gate
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_the_gate_reaches_every_member(self):
|
||||||
|
"""A gate that reaches only some members is not a weaker gate, it is no
|
||||||
|
gate: the ungated end relocates its own pages under the very transfer
|
||||||
|
the gate was installed for.
|
||||||
|
"""
|
||||||
|
for name, expected in self._EXPECTED_COVERAGE.items():
|
||||||
|
with self.subTest(composite=name):
|
||||||
|
self.assertEqual(
|
||||||
|
self._members_reached(name, "disagg_move_gate"),
|
||||||
|
expected,
|
||||||
|
f"{name}.disagg_move_gate does not cover every member",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_gate_setters_do_not_enumerate_members_themselves(self):
|
||||||
|
"""The structural half of the rule above: a setter that names its
|
||||||
|
members is one a new member silently escapes. Installation must go
|
||||||
|
through the shared helper, which drives off `_move_gate_targets`.
|
||||||
|
"""
|
||||||
|
import inspect
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.allocator import unified_hybrid_swa, unified_mamba
|
||||||
|
|
||||||
|
for name in self._EXPECTED_COVERAGE:
|
||||||
|
cls = getattr(unified_mamba, name, None) or getattr(
|
||||||
|
unified_hybrid_swa, name
|
||||||
|
)
|
||||||
|
if "set_disagg_move_gate" not in vars(cls):
|
||||||
|
continue # inherited, and the inherited one is checked above
|
||||||
|
with self.subTest(composite=name):
|
||||||
|
body = inspect.getsource(cls.set_disagg_move_gate)
|
||||||
|
self.assertIn("install_move_gate", body)
|
||||||
|
self.assertNotIn("_move_gate = ", body)
|
||||||
|
|
||||||
|
def test_swa_composite_translates_the_swa_side_separately(self):
|
||||||
|
"""The SWA sub-pool runs its OWN compaction, so a full-side physical id
|
||||||
|
does not name the SWA page holding the same virtual token. The read-path
|
||||||
|
`translate_loc_from_full_to_swa` cannot stand in either: it returns
|
||||||
|
kernel-facing ids, and the transfer addresses raw page envelopes."""
|
||||||
|
self.assertIn(
|
||||||
|
"translate_swa_indices_for_transfer",
|
||||||
|
self._own_methods("UnifiedSWATokenToKVPoolAllocator"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestEverySwaAllocatorAnswersTheTransferTranslate(CustomTestCase):
|
||||||
|
"""Any allocator with a full->SWA read translate needs the transfer sibling.
|
||||||
|
|
||||||
|
`_swa_payload` on both PD sides calls
|
||||||
|
`translate_swa_indices_for_transfer` on whatever allocator the scheduler
|
||||||
|
holds. Most get it by inheriting `SWATokenToKVPoolAllocator`, but a
|
||||||
|
composite that merely DELEGATES the read translate (the DSV4 HiSparse
|
||||||
|
allocator derives from `BaseTokenToKVPoolAllocator`) inherits neither the
|
||||||
|
default nor an override, and PD aborts with an AttributeError the moment a
|
||||||
|
sliding-window payload is built.
|
||||||
|
|
||||||
|
Derived from the live class tree rather than a hand-kept list: a list would
|
||||||
|
pass forever the day someone adds the next delegating composite.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _allocator_classes():
|
||||||
|
import importlib
|
||||||
|
import inspect
|
||||||
|
import pkgutil
|
||||||
|
|
||||||
|
import sglang.srt.mem_cache.allocator as pkg
|
||||||
|
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
|
||||||
|
|
||||||
|
found = {}
|
||||||
|
for mod_info in pkgutil.iter_modules(pkg.__path__):
|
||||||
|
try:
|
||||||
|
mod = importlib.import_module(
|
||||||
|
f"sglang.srt.mem_cache.allocator.{mod_info.name}"
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
continue # optional backends need hardware this runner may lack
|
||||||
|
for _, cls in inspect.getmembers(mod, inspect.isclass):
|
||||||
|
if issubclass(cls, BaseTokenToKVPoolAllocator):
|
||||||
|
found[cls.__name__] = cls
|
||||||
|
return found
|
||||||
|
|
||||||
|
def test_read_translate_implies_transfer_translate(self):
|
||||||
|
classes = self._allocator_classes()
|
||||||
|
# Guard the guard: an import failure that empties this set would make
|
||||||
|
# the assertion below vacuous.
|
||||||
|
self.assertIn("SWATokenToKVPoolAllocator", classes)
|
||||||
|
for name, cls in sorted(classes.items()):
|
||||||
|
if not hasattr(cls, "translate_loc_from_full_to_swa"):
|
||||||
|
continue
|
||||||
|
with self.subTest(allocator=name):
|
||||||
|
self.assertTrue(
|
||||||
|
hasattr(cls, "translate_swa_indices_for_transfer"),
|
||||||
|
f"{name} translates full->SWA for reads but cannot answer "
|
||||||
|
"translate_swa_indices_for_transfer; PD's _swa_payload "
|
||||||
|
"calls it on whatever allocator the scheduler holds",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -18,9 +18,11 @@ import unittest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.mem_cache.layout.page_major import (
|
from sglang.srt.mem_cache.layout.page_major import (
|
||||||
|
build_mha_views,
|
||||||
build_mla_views,
|
build_mla_views,
|
||||||
build_page_major_mamba_views,
|
build_page_major_mamba_views,
|
||||||
mamba_entry_bytes,
|
mamba_entry_bytes,
|
||||||
|
mha_entry_bytes,
|
||||||
mla_entry_bytes,
|
mla_entry_bytes,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -74,6 +76,118 @@ class TestMLAEnvelopeTransferAddressing(CustomTestCase):
|
|||||||
self.assertTrue(torch.equal(got, val), (page, layer, off))
|
self.assertTrue(torch.equal(got, val), (page, layer, off))
|
||||||
|
|
||||||
|
|
||||||
|
class TestMHAEnvelopeTransferAddressing(CustomTestCase):
|
||||||
|
"""The MHA counterpart of the MLA case above.
|
||||||
|
|
||||||
|
An MHA page envelope holds ``2 * layer_num`` row-blocks (layer l's K at
|
||||||
|
block 2l, its V at 2l+1). PD ships that whole envelope as one item, so a
|
||||||
|
row written through ANY per-layer view must land inside its own page's
|
||||||
|
``page_envelope_bytes`` block -- otherwise the transfer would carry a
|
||||||
|
page's K but another page's V and every kernel would still read fine
|
||||||
|
locally.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def test_page_envelope_matches_per_layer_views(self):
|
||||||
|
layer_num, page_size, head_num, head_dim, num_pages = 3, 4, 2, 8, 6
|
||||||
|
store_dtype = torch.bfloat16
|
||||||
|
entry_bytes = mha_entry_bytes(
|
||||||
|
layer_num=layer_num,
|
||||||
|
head_num=head_num,
|
||||||
|
head_dim=head_dim,
|
||||||
|
v_head_dim=head_dim,
|
||||||
|
itemsize=store_dtype.itemsize,
|
||||||
|
)
|
||||||
|
page_bytes = page_size * entry_bytes
|
||||||
|
row_bytes = head_num * head_dim * store_dtype.itemsize
|
||||||
|
self.assertEqual(page_bytes, page_size * 2 * layer_num * row_bytes)
|
||||||
|
|
||||||
|
# One page envelope of tail pad, as UnifiedKVPool allocates for MHA.
|
||||||
|
raw = torch.zeros((num_pages + 1) * page_bytes, dtype=torch.uint8)
|
||||||
|
k_views, v_views = build_mha_views(
|
||||||
|
raw,
|
||||||
|
layer_num=layer_num,
|
||||||
|
head_num=head_num,
|
||||||
|
head_dim=head_dim,
|
||||||
|
v_head_dim=head_dim,
|
||||||
|
store_dtype=store_dtype,
|
||||||
|
page_size=page_size,
|
||||||
|
num_pages=num_pages,
|
||||||
|
anchor_bytes=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
blocks = 2 * layer_num
|
||||||
|
for page in range(num_pages):
|
||||||
|
for layer in range(layer_num):
|
||||||
|
for is_v, views in ((0, k_views), (1, v_views)):
|
||||||
|
for pos in range(page_size):
|
||||||
|
row = page * blocks * page_size + pos
|
||||||
|
views[layer][row].fill_(1)
|
||||||
|
(nz,) = torch.nonzero(raw, as_tuple=True)
|
||||||
|
lo, hi = int(nz.min()), int(nz.max())
|
||||||
|
self.assertGreaterEqual(
|
||||||
|
lo,
|
||||||
|
page * page_bytes,
|
||||||
|
f"page={page} layer={layer} v={is_v} pos={pos} "
|
||||||
|
"wrote below its page envelope",
|
||||||
|
)
|
||||||
|
self.assertLess(
|
||||||
|
hi,
|
||||||
|
(page + 1) * page_bytes,
|
||||||
|
f"page={page} layer={layer} v={is_v} pos={pos} "
|
||||||
|
"wrote past its page envelope",
|
||||||
|
)
|
||||||
|
views[layer][row].zero_()
|
||||||
|
|
||||||
|
def test_envelope_move_is_a_whole_page_copy(self):
|
||||||
|
"""Relocating a page envelope must move every layer's K and V with it;
|
||||||
|
this is what `UnifiedMHATokenToKVPool.move_kv_cache` relies on and what
|
||||||
|
makes a physical page id a valid PD transfer index after compaction."""
|
||||||
|
layer_num, page_size, head_num, head_dim, num_pages = 2, 2, 1, 4, 4
|
||||||
|
store_dtype = torch.bfloat16
|
||||||
|
entry_bytes = mha_entry_bytes(
|
||||||
|
layer_num=layer_num,
|
||||||
|
head_num=head_num,
|
||||||
|
head_dim=head_dim,
|
||||||
|
v_head_dim=head_dim,
|
||||||
|
itemsize=store_dtype.itemsize,
|
||||||
|
)
|
||||||
|
page_bytes = page_size * entry_bytes
|
||||||
|
raw = torch.zeros((num_pages + 1) * page_bytes, dtype=torch.uint8)
|
||||||
|
k_views, v_views = build_mha_views(
|
||||||
|
raw,
|
||||||
|
layer_num=layer_num,
|
||||||
|
head_num=head_num,
|
||||||
|
head_dim=head_dim,
|
||||||
|
v_head_dim=head_dim,
|
||||||
|
store_dtype=store_dtype,
|
||||||
|
page_size=page_size,
|
||||||
|
num_pages=num_pages,
|
||||||
|
anchor_bytes=0,
|
||||||
|
)
|
||||||
|
blocks = 2 * layer_num
|
||||||
|
# Distinct content in source page 1, every layer, K and V.
|
||||||
|
for layer in range(layer_num):
|
||||||
|
for pos in range(page_size):
|
||||||
|
row = 1 * blocks * page_size + pos
|
||||||
|
k_views[layer][row].fill_(layer + 1)
|
||||||
|
v_views[layer][row].fill_(-(layer + 1))
|
||||||
|
|
||||||
|
env = raw[: num_pages * page_bytes].view(num_pages, page_bytes)
|
||||||
|
env[3] = env[1]
|
||||||
|
|
||||||
|
for layer in range(layer_num):
|
||||||
|
for pos in range(page_size):
|
||||||
|
row = 3 * blocks * page_size + pos
|
||||||
|
self.assertTrue(
|
||||||
|
torch.all(k_views[layer][row] == layer + 1),
|
||||||
|
f"K layer {layer} did not ride the envelope move",
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.all(v_views[layer][row] == -(layer + 1)),
|
||||||
|
f"V layer {layer} did not ride the envelope move",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestMambaEnvelopeTransferAddressing(CustomTestCase):
|
class TestMambaEnvelopeTransferAddressing(CustomTestCase):
|
||||||
def test_slot_envelope_is_self_contained(self):
|
def test_slot_envelope_is_self_contained(self):
|
||||||
"""A slot's conv+temporal state for all layers must live exactly in
|
"""A slot's conv+temporal state for all layers must live exactly in
|
||||||
|
|||||||
@@ -345,11 +345,12 @@ class TestUnifiedMHATokenToKVPool(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def test_transfer_entry_points_fail_loud(self):
|
def test_transfer_entry_points_fail_loud(self):
|
||||||
"""PD / CPU-copy entry points assume per-layer buffers indexed by TOKEN
|
"""The entry points that assume per-layer buffers indexed by TOKEN id
|
||||||
id and would silently mis-index the row space, so each must raise."""
|
would silently mis-index against the row space (or hit a missing-attr
|
||||||
|
AttributeError), so each must raise. `get_contiguous_buf_infos` is NOT
|
||||||
|
among them: PD addresses this pool as whole page envelopes, pinned by
|
||||||
|
`test_pd_registration_is_one_whole_envelope` below."""
|
||||||
_, pool = _make_pool_and_kv(1)
|
_, pool = _make_pool_and_kv(1)
|
||||||
with self.assertRaises(NotImplementedError):
|
|
||||||
pool.get_contiguous_buf_infos()
|
|
||||||
with self.assertRaises(NotImplementedError):
|
with self.assertRaises(NotImplementedError):
|
||||||
pool.get_cpu_copy(torch.tensor([1]))
|
pool.get_cpu_copy(torch.tensor([1]))
|
||||||
with self.assertRaises(NotImplementedError):
|
with self.assertRaises(NotImplementedError):
|
||||||
@@ -357,6 +358,24 @@ class TestUnifiedMHATokenToKVPool(unittest.TestCase):
|
|||||||
with self.assertRaises(NotImplementedError):
|
with self.assertRaises(NotImplementedError):
|
||||||
pool.set_kv_buffer_prefix_valid()
|
pool.set_kv_buffer_prefix_valid()
|
||||||
|
|
||||||
|
def test_pd_registration_is_one_whole_envelope(self):
|
||||||
|
"""PD registers ONE region -- the whole raw buffer -- with the page
|
||||||
|
envelope as the item, so the transfer engine addresses it as
|
||||||
|
`raw_ptr + physical_page * page_envelope_bytes`. Per-layer regions
|
||||||
|
would be wrong here: the per-layer views overlap inside the envelope
|
||||||
|
and index in kernel-facing ids, not token ids."""
|
||||||
|
kv, pool = _make_pool_and_kv(1)
|
||||||
|
ptrs, lens, item_lens = pool.get_contiguous_buf_infos()
|
||||||
|
self.assertEqual(len(ptrs), 1)
|
||||||
|
self.assertEqual(len(lens), 1)
|
||||||
|
self.assertEqual(len(item_lens), 1)
|
||||||
|
self.assertEqual(ptrs[0], kv._raw.data_ptr())
|
||||||
|
self.assertEqual(lens[0], kv._raw.numel())
|
||||||
|
self.assertEqual(item_lens[0], pool._page_bytes)
|
||||||
|
# The whole addressable page range must fit the registered region, or
|
||||||
|
# the last page's write would run off the end of the RDMA mapping.
|
||||||
|
self.assertLessEqual(pool._num_pages * item_lens[0], lens[0])
|
||||||
|
|
||||||
def test_hnd_env_cannot_hijack_layout(self):
|
def test_hnd_env_cannot_hijack_layout(self):
|
||||||
"""SGLANG_USE_HND_KVCACHE must not flip this pool's layout: HND indexes
|
"""SGLANG_USE_HND_KVCACHE must not flip this pool's layout: HND indexes
|
||||||
4-D while the per-layer views are 3-D, so the pinned label has to win."""
|
4-D while the per-layer views are 3-D, so the pinned label has to win."""
|
||||||
|
|||||||
Reference in New Issue
Block a user